OSCR

The Tower Foraging Park: A paradigm for studying cognitive and motor processes underlying behavioral flexibility in freely moving mice.

Code ↔ Paper

14 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 14 matches
  1. [1] § STAR★Methods › Quantification and statistical analysis › Statistical analyses ↔ FigureS02_group1_SpeedvsRank.ipynb, lines 100–202 · score 0.75 · linear regression, negative slope, CCW turn, shuffled, CCW QTs, rank
  2. [2] § Results › Experiment 3: Adaptation of group 1 and group 2 mice to random foraging rules ↔ Figure09left_group1_randomProtocol_QT_stats.ipynb, lines 1213–1324 · score 0.71 · protocol change, unrewarded QTs, CCW turns, individual mice, CCW QTs, speed
  3. [3] § Results › Experiment 3: Adaptation of group 1 and group 2 mice to random foraging rules ↔ Figure09right_group2_randomProtocol_QT_stats.ipynb, lines 1238–1358 · score 0.71 · protocol change, unrewarded QTs, CCW turns, individual mice, CCW QTs, speed
  4. [4] § STAR★Methods › Quantification and statistical analysis › Statistical analyses ↔ DraftPermutationBootstrap.ipynb, lines 117–235 · score 0.65 · confidence interval, 97.5 %, 2.5 %, global, permutation, medians
  5. [5] § STAR★Methods › Quantification and statistical analysis › Data processing ↔ Figure03_group1_qt_example_and_stats.ipynb, lines 1262–1388 · score 0.63 · speed threshold, acceleration, backward, initiation, crossing, duration
  6. [6] § STAR★Methods › Quantification and statistical analysis › Statistical analyses ↔ FigureS01_TestdifferencebetweenCCWvsCWkinematics.ipynb, lines 623–746 · score 0.60 · median Fr chet, randomly swapped, downsampling, Spearman, distance, trajectories
  7. [7] § Results › Experiment 1: Stable unidirectional harvesting followed by reversal in group 1 mice ↔ Figure03_group1_qt_example_and_stats.ipynb, lines 1262–1388 · score 0.58 · low speed threshold, continuous epochs, shorter, Video, movement, Figure 3
  8. [8] § STAR★Methods › Quantification and statistical analysis › Data processing ↔ Figure5MaudAlternation.ipynb, lines 117–199 · score 0.58 · speed threshold, acceleration, backward, crossing, duration, videos
  9. [9] § STAR★Methods › Quantification and statistical analysis › Data processing ↔ Figure05_group1_exploit_explore.ipynb, lines 896–1026 · score 0.57 · video frame, pixels, smoothed, SE, NE, NW
  10. [10] § Results › Experiment 1: Stable unidirectional harvesting followed by reversal in group 1 mice ↔ Figure5MaudAlternation.ipynb, lines 117–199 · score 0.56 · low speed threshold, continuous epochs, shorter, Video
  11. [11] § STAR★Methods › Quantification and statistical analysis › Statistical analyses ↔ Figure09right_group2_randomProtocol_QT_stats.ipynb, lines 1238–1358 · score 0.55 · peak speed, CCW CW, Fr chet, permutation, distance, median
  12. [12] § Results › Experiment 1: Stable unidirectional harvesting followed by reversal in group 1 mice ↔ Figure05_group1_exploit_explore.ipynb, lines 673–771 · score 0.54 · unrewarded QT, raster, visit, exploitation, Figure 5, towers
  13. [13] § STAR★Methods › Quantification and statistical analysis › Kinematic analyses ↔ Figure03_group1_qt_example_and_stats.ipynb, lines 961–1067 · score 0.52 · rotating trajectories, transformed, NW, SW, corner, QT
  14. [14] § STAR★Methods › Quantification and statistical analysis › Kinematic analyses ↔ Figure08_group2_qt_stats.ipynb, lines 1136–1240 · score 0.51 · rotating trajectories, transformed, NW, SW, corner, QT

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

Jupyter notebook · 2,541 lines · 119 KB · no license · 3 matches

  1. # %% [markdown]
  2. # # Notebook For Figure 03 Method Paper
  3. # %% [markdown]
  4. # ## Todos
  5. # %%
  6. from IPython.display import display
  7. from PIL import Image
  8. # Load and display the image
  9. img = Image.open("Figure03.png")
  10. display(img)
  11. # %% [markdown]
  12. # # 1. Importing necessary libraries and general functions
  13. # %%
  14. import os
  15. import glob
  16. from processing_TowerCoordinates import *
  17. from processing_session_trajectory import *
  18. import matplotlib.pyplot as plt
  19. from matplotlib.ticker import MaxNLocator
  20. import matplotlib.path as mpath
  21. import matplotlib.cm as cm
  22. import matplotlib.patches as patches
  23. from matplotlib.gridspec import GridSpec
  24. from matplotlib.colors import Normalize
  25. from mpl_toolkits.mplot3d import Axes3D
  26. from mpl_toolkits.axes_grid1.inset_locator import inset_axes
  27. import numpy as np
  28. import pickle
  29. import scipy.ndimage as ndimage
  30. from scipy.stats import wilcoxon
  31. from processing_TowerCoordinates import *
  32. from datetime import date
  33. from scipy.ndimage import gaussian_filter as smooth
  34. import matplotlib.colors as mcolors
  35. import similaritymeasures as sm
  36. from bisect import bisect
  37. from scipy.stats import spearmanr
  38. import warnings
  39. from matplotlib.colors import LinearSegmentedColormap
  40. plt.style.use('paper.mplstyle')
  41. # useful line to interupt Run all execution
  42. # raise SystemExit("Stopping execution here.")
  43. # %%
  44. plotintermediatesteps = False
  45. # %% [markdown]
  46. # # 2. Select mice that will be analysed in the figure
  47. # ### Define the data_folder where your MOU* folders are, enter the list of mice (MOU* folders), define the sessions that will be used for each animals
  48. # %%
  49. # defining data folder path and mice list
  50. # path_to_data_folder is the path of the folder where you store the folders of your different mice.
  51. path_to_data_folder='/LocalData/ForagingMice/4TowersTaskMethodPaper_Data/Group1Data/'
  52. path_to_data_folder = '/home/tom/Code/data/tower_foraging_park_data/Group1Data'
  53. # path_to_data_folder = '/Users/davidrobbe/Documents/Science/Data/ForagingMice/Group1Data/'
  54. # path_to_data_folder='/home/david/MyLocalData/4TowersTask_MethodPaper/Group1Data'
  55. # Analysing the entire group of mice
  56. mice_to_analyse = [
  57. "MOUEml1_5", "MOUEml1_8", "MOUEml1_11", "MOUEml1_12", "MOUEml1_13", "MOUEml1_15", "MOUEml1_18", "MOUEml1_20",
  58. "MOURhoA_2", "MOURhoA_5", "MOURhoA_6", "MOURhoA_8", "MOURhoA_9", "MOURhoA_12", "MOURhoA_14",
  59. "MOUB6NN_4", "MOUB6NN_6", "MOUB6NN_13", "MOUB6NN_15"
  60. ]
  61. # Verify that all folders in mice_to_analyse are present in path_to_data_folder
  62. missing_folders = [mouse for mouse in mice_to_analyse if not os.path.isdir(os.path.join(path_to_data_folder, mouse))]
  63. if missing_folders:
  64. print("Missing mice folders:", missing_folders)
  65. else:
  66. print("All mice folders are present in data folder.")
  67. # Print the number of mice, the list of mice
  68. print(f' {len(mice_to_analyse)} {"mice" if len(mice_to_analyse) > 1 else "mouse"} will be analysed\n')
  69. # Select the number of sessions that will be analysed for each mice knowing the analysis starts with the first session (familiarisation)
  70. first_and_last_session_indexes = [0,14]
  71. # Setting the seed for random processes used in statistics
  72. seed = 27
  73. np.random.seed(seed)
  74. # Select the offset to apply to the number of the sessions when plotting. By default, it is equal to the number of the first session.
  75. # It resulsts in sessions being named respectively to their actual positions in the data set. TODO: Should reformulate this
  76. session_index_offset = first_and_last_session_indexes[0]
  77. # %% [markdown]
  78. # # 3. Chosing illustration mouse and sessions
  79. # %%
  80. # Chose the 2 mice that are highlighted in statistics across figures
  81. illustration_mice = ['MOUEml1_8', 'MOUEml1_5']
  82. # Chose the color with which the illustration mice will highlighted with
  83. illustration_colors = ['darkorange', 'green']
  84. # Mouse that will be used as single example in this figure to show QTs trajectories and kinematics
  85. example_mouse_index = 0
  86. example_mouse_bis_index = 1
  87. example_mouse = illustration_mice[example_mouse_index]
  88. example_mouse_bis = illustration_mice[example_mouse_bis_index]
  89. # Chosing index of the sessions to show as examples. They are chosen by their index (session 0 is the first session etc.)
  90. illustration_sessions_indexes = [0, 5, 12]
  91. # coordinates in cm of the external walls of the arena
  92. arena_coordinates_cm = [[4.5, 88.86], [90.3, 88.86], [90.3, 2.7], [4.5, 2.7]]
  93. # Set the limits of the x-axis on the statistics plots
  94. xlim_stats = [first_and_last_session_indexes[0]+0.5,first_and_last_session_indexes[1]+0.5]
  95. # %% [markdown]
  96. # # 4. General functions
  97. # %%
  98. def get_day_and_period(session_idx):
  99. day = session_idx // 2 + 1
  100. period = 'AM' if session_idx % 2 == 0 else 'PM'
  101. return day, period
  102. def force_aspect(ax, ratio=1):
  103. """
  104. Force the aspect ratio of the given axis (ax) to a specific ratio.
  105. The ratio parameter allows scaling of the aspect ratio. Default is 1.
  106. Arguments:
  107. ax (matplotlib.axes.Axes): Matplotlib axis object for plotting.
  108. ratio (float, optional): Ratio of the figure's dimensions
  109. """
  110. ratio = 1.0 # Set ratio to 1.0, making the aspect ratio 1:1 by default
  111. # Get the current limits of the x and y axes
  112. x_left, x_right = ax.get_xlim() # Get the left and right x-axis limits
  113. y_low, y_high = ax.get_ylim() # Get the lower and upper y-axis limits
  114. # Calculate the new aspect ratio and set it
  115. # The formula is the absolute ratio of the width and height, adjusted by the given ratio
  116. ax.set_aspect(abs((x_right - x_left) / (y_low - y_high)) * ratio)
  117. def cm2inch(value):
  118. """
  119. Converts centimeters to inches for figure size.
  120. Arguments:
  121. value (float): value to convert in cm
  122. """
  123. return value/2.54
  124. def filter_qts(qts):
  125. """
  126. Filter out QTs with length above threshold and during which more than 1 switch occured.
  127. Arguments:
  128. qts (list): list of QTs (in the original format from the pickle)
  129. return:
  130. (list): list of filtered QTs
  131. """
  132. filtered_qts = []
  133. for qt in qts:
  134. # Skip current QT if there is more than 1 switches
  135. if qt[3]['num_trapezeswitch']!=1:
  136. continue
  137. # Skip current QT if it is longer than a threshold length (30 cm)
  138. if qt[4]['epoch_distance']>30:
  139. continue
  140. # Keep current QT if it did not check the removal conditions
  141. filtered_qts.append(qt)
  142. return filtered_qts
  143. def count_qts(qts):
  144. total_qts = len(qts)
  145. cw_qts = 0
  146. rewarded_qts = 0
  147. for qt in qts:
  148. if qt[3]['direction']=='CW':
  149. cw_qts += 1
  150. if qt[3]['Rewarded']:
  151. rewarded_qts += 1
  152. return total_qts, cw_qts, rewarded_qts
  153. def compute_average_direction(points):
  154. """
  155. Compute the average direction of a trajectory portion, assuming the coordinates are in chronological order.
  156. Arguments:
  157. points (2D numpy.array): Trajectory of which the average direction will be computed
  158. Outputs:
  159. avg_diff (1D numpy.array): vector pointing to the average direction
  160. angle (float): angle of the vector with the x-axis
  161. """
  162. diffs = np.diff(points, axis=0)
  163. avg_diff = np.mean(diffs, axis=0)
  164. angle = np.arctan2(avg_diff[1], avg_diff[0])
  165. return avg_diff, angle
  166. def finding_mouse_rewarded_direction(folder_path_mouse_to_process, session_index):
  167. """
  168. Determines the rewarded direction for the session corresponding to the input index of a given mouse.
  169. This index is relative to the session position in the series of analysed sessions.
  170. This input index can have an offset. This is usefull if the sessions series analysed does not start with the first session.
  171. Arguments:
  172. folder_path_mouse_to_process (str): Path to the folder containing mouse sessions folders.
  173. session_index (int): Index of the session that will be used to define the rewarded direction
  174. Returns:
  175. str: 'CW' (Clockwise) if the rewarded direction is 270 degrees,
  176. 'CCW' (Counterclockwise) if the rewarded direction is 90 degrees,
  177. numpy.nan if reward delivery is not allowed
  178. None if an error occurs or if both directions are rewarded.
  179. """
  180. # Get all session folders that start with 'MOU' and sort them
  181. sessions_to_process = sorted([name for name in os.listdir(folder_path_mouse_to_process)
  182. if os.path.isdir(os.path.join(folder_path_mouse_to_process, name))
  183. and name.startswith('MOU')])
  184. # Load data from the last session
  185. session_traj_df, session_turns_df, session_param_df = load_data(folder_path_mouse_to_process, sessions_to_process[session_index])
  186. # Extract rewarded direction in degrees
  187. rewarded_direction_degrees = session_param_df["potentialRewardedDirections"][0]
  188. # Check if reward delivery is allowed
  189. if session_param_df["allowRewardDelivery"][0]:
  190. # Determine the rewarded direction based on the extracted value
  191. if rewarded_direction_degrees == '[270]':
  192. rewarded_direction = 'CW' # Clockwise
  193. elif rewarded_direction_degrees == '[90]':
  194. rewarded_direction = 'CCW' # Counterclockwise
  195. elif rewarded_direction_degrees == '[90, 270]':
  196. rewarded_direction = 'both' # Clockwise and Counterclockwise
  197. # Returns None if the rewarded direction entry in session_param_df is not recognized
  198. else:
  199. print('ERROR: Unexpected rewarded direction value:', rewarded_direction_degrees)
  200. return None # Explicitly return None to indicate failure
  201. else:
  202. # Rewarded direction is set to X if reward delivery os not allowed
  203. rewarded_direction = 'X'
  204. return rewarded_direction
  205. def plot_learning_curves(mouse_metric_persession, ax, mice_list=None, mice_to_highlight=[], mice_to_highlight_labels=[None,None], highlight_colors = ["darkorange", "green"], show_individual_mice=True,
  206. median_color='black', show_xlabel=True, ylabel='', main_line_label=None, show_fam=False,
  207. tick_interval=1, index_offset=0, xlim=[None,None], ylim=None, show_legend=True, legend_loc=(0.05, 0.9)):
  208. """
  209. Plots a given metric across sessions for multiple mice (also called a learning curves).
  210. Arguments:
  211. mouse_metric_persession (dict): Dictionary where keys are mouse IDs and values are lists of (session_index, value) lists.
  212. mice_list (list, optional): List of mice to include in the plot. Defaults to all mice in the mouse_metric_persession.
  213. mice_to_highlight (list, optional): List of up to two mice IDs to highlight with distinct colors.
  214. mice_to_highlight_labels (list, optional): List of the label of the mice to highlight.
  215. show_individual_mice (bool): If true, plot a line for each individual mouse.
  216. ax (matplotlib.axes.Axes): Axes object for plotting.
  217. show_xlabel (bool, optional): Whether to display the x-axis label. Defaults to True.
  218. ylabel (str, optional): Label for the y-axis.
  219. main_line_label (str, optional): Label of the median and quartiles range.
  220. show_fam (bool): Toggle the display of familiarization sessions label
  221. tick_interval (int, optional): Interval for x-axis tick marks. Defaults to 1.
  222. index_offset (int, optional): Offset to apply to the numbers in the x axis. Use it if you want to the first sessions to be named '1'.
  223. xlim (tuple, optional): Tuple specifying x-axis limits. Defaults to None.
  224. ylim (tuple, optional): Tuple specifying y-axis limits. Defaults to None.
  225. show_legend (bool, optional): Whether to show the legend. Defaults to True.
  226. legend_loc (tupple, optional): Location of the legend.
  227. """
  228. # If no mice are specified, use all available mice
  229. if mice_list is None:
  230. mice_list = list(mouse_metric_persession.keys())
  231. all_session_indices = set() # Track all session indices across selected mice
  232. values_per_session = {} # Store values for each session across mice
  233. mice_to_plot = copy.deepcopy(mice_list)
  234. # Ensure that the mice to hihlight will be plotted and displayed on the legend in order
  235. if len(mice_to_highlight)>0:
  236. mice_to_plot.remove(mice_to_highlight[0])
  237. mice_to_plot.remove(mice_to_highlight[1])
  238. mice_to_plot = mice_to_highlight + mice_to_plot
  239. # Loop through each mouse and gather session data
  240. for mouse in mice_to_plot:
  241. if mouse not in mouse_metric_persession:
  242. print(f"Mouse {mouse} not found in data. Skipping.")
  243. continue
  244. # Selecting the sub-list of sessions to analyse
  245. sessions = mouse_metric_persession[mouse]
  246. # Extract session indices and corresponding values
  247. session_indices = [session[0] for session in sessions]
  248. values = [session[1] for session in sessions]
  249. # Check if there is a nan value in the values list and remove it to allow plotting of lines between points with values
  250. if sum(np.isnan(x) for x in values)>0:
  251. session_indices_array=np.array(session_indices)
  252. values_array=np.array(values)
  253. mask = ~np.isnan(values_array)
  254. session_indices=session_indices_array[mask]
  255. values=values_array[mask]
  256. # Plot learning curves for each mouse if show_individual_mice is True
  257. if show_individual_mice:
  258. if mouse == mice_to_highlight[0]:
  259. color = highlight_colors[0]
  260. ax.plot(np.array(session_indices)+index_offset, values, color=color, marker='o', linewidth=0.7,
  261. markersize=1, label= mice_to_highlight_labels[0], zorder=100*len(median_color))
  262. elif mouse == mice_to_highlight[1]:
  263. color = highlight_colors[1]
  264. ax.plot(np.array(session_indices)+index_offset, values, color=color, marker='o', linewidth=0.7,
  265. markersize=1, label= mice_to_highlight_labels[1], zorder=100*len(median_color))
  266. else:
  267. ax.plot(np.array(session_indices)+index_offset, values, marker='o', linewidth=0.7, markersize=1, alpha=0.3, markeredgewidth=0.0)
  268. # Update session index tracking
  269. all_session_indices.update(session_indices)
  270. for session, value in sessions:
  271. if session not in values_per_session:
  272. values_per_session[session] = []
  273. values_per_session[session].append(value)
  274. # Convert session indices to a sorted list
  275. sorted_sessions = sorted(all_session_indices)
  276. # Compute median and 25th-75th percentile range for each session
  277. median_values = []
  278. lower_quartile = []
  279. upper_quartile = []
  280. for session in sorted_sessions:
  281. session_values = values_per_session[session]
  282. median_values.append(np.nanmedian(session_values))
  283. lower_quartile.append(np.nanpercentile(session_values, 25))
  284. upper_quartile.append(np.nanpercentile(session_values, 75))
  285. median_values = np.array(median_values)
  286. lower_quartile = np.array(lower_quartile)
  287. upper_quartile = np.array(upper_quartile)
  288. # If the median and quartiles range don't have a label, give them a default label.
  289. if not main_line_label:
  290. main_line_label = f'Median \u00B1 Quartiles, n={len(mice_list)}'
  291. # Plot median learning curve with shaded 25th-75th percentile range
  292. ax.errorbar(np.array(sorted_sessions)+index_offset, median_values, yerr=[median_values-lower_quartile, upper_quartile-median_values], alpha=1, color=median_color, linewidth = 2, label=main_line_label, zorder=50*len(median_color))
  293. # Set axis labels and formatting
  294. if show_xlabel:
  295. ax.set_xlabel('Session Number', fontsize=7)
  296. ax.set_ylabel(ylabel, fontsize=7)
  297. # Ensure x-axis labels are integers
  298. ax.xaxis.set_major_locator(MaxNLocator(integer=True))
  299. # Set x-ticks at specified intervals
  300. if all_session_indices:
  301. max_index = max(all_session_indices)
  302. ax.set_xticks(range(1+index_offset, max_index + 1+index_offset, tick_interval))
  303. # Apply x and y axis limits if provided
  304. if xlim is not None:
  305. ax.set_xlim(xlim)
  306. if ylim is not None:
  307. ax.set_ylim(ylim)
  308. ylimits = ax.get_ylim() # Get the current y-limits
  309. rectangle_height = abs(ylimits[1]-ylimits[0]) # Calculate the height of the rectangle
  310. # Define the rectangle position and size
  311. rect = patches.Rectangle((0.5, ylimits[0]), 2, rectangle_height, color='moccasin', alpha=1, lw=0, zorder=-np.inf) # (x_start, y_start), width, height
  312. # Add rectangle to the plot
  313. ax.add_patch(rect)
  314. if show_fam:
  315. # Add text label on top of the rectangle
  316. ax.text(1.5, ylimits[1], "Fam.", fontsize=6, ha='center', va='bottom', fontweight='normal',color='orange', alpha=0.6)
  317. # Show legend if required
  318. if show_legend:
  319. ax.legend(frameon=False, loc=legend_loc, )
  320. def shuffled_spearman_test(mouse_values_persession, shuffle_number, hypothesis, illustration_mouse_name = None, first_and_last_session_indexes=[None,None]):
  321. """
  322. Compute the Spearman correlation coefficient of input session-wise data then compare it to a null distribution computed by shuffling
  323. the input data several times
  324. Arguments:
  325. mouse_values_persession (dict): Dictionary with session numbers as keys and the corresponding values as entries
  326. shuffle_number (int): Number of time the data will be shuffled and the corresponding correlation coefficient computed
  327. hypothesis (string): Hypothesis tested, whether the values sequence is increasing or decreasing across sessions
  328. illustration_mouse_name (list, optional): List of mice from mouse_values_persession to use
  329. first_and_last_session_indexes (list, optional): List containing indexes of first and last session to plot
  330. Return:
  331. (list): list containing the Spearman correlation coefficient of the unshuffled data and the p-value associaed to the test
  332. """
  333. # Selecting mice to analyse
  334. if illustration_mouse_name is None:
  335. illustration_mouse_name = list(mouse_values_persession.keys())
  336. ### Compute the Spearman correlation coefficient distribution by shuffling data ###
  337. # Initialize the list where the Spearman correlation coefficients from shuffled data will be stored.
  338. # This will be the null distribution of Spearman correlation coefficient
  339. spearman_rho_list = []
  340. # Shuffle the data shuffle_number times
  341. for _ in range(shuffle_number):
  342. # Initialize the list that will contain shuffled data
  343. shuffled_values_list = []
  344. # Iterate on every mice
  345. for mouse in illustration_mouse_name:
  346. # Checking if the current mouse has no value to plot
  347. if mouse not in mouse_values_persession:
  348. print(f"Mouse {mouse} not found in data. Skipping.")
  349. continue
  350. # Selecting sessions under the limit
  351. sessions = copy.deepcopy(mouse_values_persession[mouse][first_and_last_session_indexes[0]:first_and_last_session_indexes[1]])
  352. # Extract sessions number
  353. session_indices = [session[0] for session in sessions]
  354. # Extract and shuffle values
  355. shuffled_values = np.random.choice([session[1] for session in sessions], size=len(sessions), replace=False)
  356. # Store the shuffled values
  357. shuffled_values_list.append(shuffled_values)
  358. # Compute the median over all the mice
  359. median_values = np.nanmedian(shuffled_values_list,axis=0)
  360. # Compute the Spearman correlation coefficient of the median
  361. spearman_result = spearmanr(session_indices, median_values)
  362. # Store the coefficient
  363. spearman_rho_list.append(spearman_result[0])
  364. ### Compute the Spearman correlation coefficient of unshuffled data ###
  365. values_list = []
  366. for mouse in illustration_mouse_name:
  367. # Checking if the current mouse has no value to plot
  368. if mouse not in mouse_values_persession:
  369. print(f"Mouse {mouse} not found in data. Skipping.")
  370. continue
  371. # Selecting sessions under the limit
  372. sessions = copy.deepcopy(mouse_values_persession[mouse][first_and_last_session_indexes[0]:first_and_last_session_indexes[1]])
  373. # Extract sessions number
  374. session_indices = [session[0] for session in sessions]
  375. # Extract values
  376. values = [session[1] for session in sessions]
  377. # Store the value
  378. values_list.append(values)
  379. # Compute the median over all the mice
  380. actual_median_values = np.nanmedian(values_list,axis=0)
  381. # Compute the Spearman correlation coefficient of the median
  382. actual_spearman_result = spearmanr(session_indices,actual_median_values, nan_policy='omit')
  383. actual_rho = actual_spearman_result[0]
  384. # Compute the p-value corresponding to the tested hypothesis
  385. if hypothesis=="increasing":
  386. # Computing the the proportion of the null distribution above the correlation coefficient of unshuffled data. This is the p-value
  387. p_value = len(spearman_rho_list[bisect(np.sort(spearman_rho_list), actual_rho):])/len(spearman_rho_list)
  388. elif hypothesis=="decreasing":
  389. # Computing the the proportion of the null distribution below the correlation coefficient of unshuffled data. This is the p-value
  390. p_value = len(spearman_rho_list[:bisect(np.sort(spearman_rho_list), actual_rho)])/len(spearman_rho_list)
  391. else:
  392. # Raise an error if no valid hypothesis was made
  393. print("ERROR : Invalid hypothesis. Please specify increasing or decreasing")
  394. return
  395. return actual_rho, p_value
  396. def plot_shuffled_spearman_test_res(ax, text_loc, mouse_values_persession, n_shuffles, hypothesis, color='black', illustration_mouse_name=None, first_and_last_session_indexes=[1,None]):
  397. """
  398. Plots the result of shuffled_spearman_test as a text on a figure.
  399. Arguments:
  400. ax (matplotlib.axes.Axes): Axes object for plotting
  401. text_loc (list): coordinate of the text (in data units)
  402. mouse_values_persession (dict): Dictionary with session numbers as keys and the corresponding values as entries
  403. shuffle_number (int): Number of time the data will be shuffled and the corresponding correlation coefficient computed
  404. hypothesis (string): Hypothesis tested, whether the values sequence is increasing or decreasing across sessions
  405. color (str): Color of the text
  406. illustration_mouse_name (list, optional): List of mice from mouse_values_persession to use
  407. first_and_last_session_indexes (list, optional): List containing indexes of first and last session to plot
  408. """
  409. test_res = shuffled_spearman_test(mouse_values_persession, n_shuffles, hypothesis, illustration_mouse_name=illustration_mouse_name, first_and_last_session_indexes=first_and_last_session_indexes)
  410. if test_res[1]==0:
  411. ax.text(text_loc[0],text_loc[1],rf'$\rho = {round(test_res[0],2)}$' + '\n' + rf'p$ < 0.001$', fontsize=5, color=color)
  412. else:
  413. ax.text(text_loc[0],text_loc[1],rf'$\rho = {round(test_res[0],2)}$' + '\n' + rf'p$ = {round(test_res[1],3)}$', fontsize=5, color=color)
  414. # %% [markdown]
  415. # # 5. Computations for panel B, D and J statistics
  416. # %% [markdown]
  417. # The next two functions are used to:
  418. # - Find the times at which QTs occured.
  419. # - Find the time at which the number of QTs becomes higher than a given threshold. (not clear)
  420. # %%
  421. def cumulated_turns_time_profile(folder_path_mouse_to_process, session_to_process, rewarded_direction):
  422. """
  423. Computes the cumulative number of runs around a tower at each time a new run around a tower occurs, for a given session and a given mouse.
  424. Arguments:
  425. folder_path_mouse_to_process (str): Path to the folder containing mouse session data.
  426. session_to_process (str): Name of the session to be processed.
  427. reference_session_index (int): Index of the session that will be used to define the rewarded direction
  428. Returns:
  429. tuple: Two lists containing:
  430. - (good_turns_time, cumulated_good_turns): Sorted times and cumulative counts for turns in rewarding direction.
  431. - (bad_turns_time, cumulated_bad_turns): Sorted times and cumulative counts for turns in non-rewarded direction.
  432. """
  433. # Construct the path to the session pickle file
  434. output_pickle_filename = f"{session_to_process}_basic_processing_output.pickle"
  435. output_pickle_filepath = os.path.join(folder_path_mouse_to_process, session_to_process, output_pickle_filename)
  436. # Load session data from the pickle file
  437. with open(output_pickle_filepath, 'rb') as file:
  438. session_data = pickle.load(file)
  439. # Deep copy the list of turn epochs
  440. runs_around_tower = filter_qts(copy.deepcopy(session_data['all_epochs']['run_around_tower']))
  441. # Initialize lists to store filtered turns
  442. time_of_runsaroundtower_cw = []
  443. time_of_runsaroundtower_ccw = []
  444. # Iterate through each recorded turn around the tower
  445. for run_around_tower in runs_around_tower:
  446. # Categorize turns based on direction
  447. if run_around_tower[3]['direction'] == 'CW':
  448. time_of_runsaroundtower_cw.append(run_around_tower[4]['epoch_time'])
  449. elif run_around_tower[3]['direction'] == 'CCW':
  450. time_of_runsaroundtower_ccw.append(run_around_tower[4]['epoch_time'])
  451. # Sort the turn times for cumulative calculations
  452. CW_times_sorted = np.sort(time_of_runsaroundtower_cw)
  453. CCW_times_sorted = np.sort(time_of_runsaroundtower_ccw)
  454. # Compute cumulative counts for each direction
  455. CW_cumulative = np.arange(1, len(CW_times_sorted) + 1)
  456. CCW_cumulative = np.arange(1, len(CCW_times_sorted) + 1)
  457. # Assign good and bad turns based on the rewarded direction
  458. if rewarded_direction == 'CW':
  459. good_turns_time = CW_times_sorted
  460. bad_turns_time = CCW_times_sorted
  461. cumulated_good_turns = CW_cumulative
  462. cumulated_bad_turns = CCW_cumulative
  463. elif rewarded_direction == 'CCW':
  464. good_turns_time = CCW_times_sorted
  465. bad_turns_time = CW_times_sorted
  466. cumulated_good_turns = CCW_cumulative
  467. cumulated_bad_turns = CW_cumulative
  468. elif rewarded_direction == 'X':
  469. good_turns_time = np.ones(len(CCW_times_sorted))*np.nan
  470. bad_turns_time = np.ones(len(CW_times_sorted))*np.nan
  471. cumulated_good_turns = np.ones(len(CCW_cumulative))*np.nan
  472. cumulated_bad_turns = np.ones(len(CW_cumulative))*np.nan
  473. else:
  474. print('ERROR: Unexpected rewarded direction value')
  475. return None # Explicitly return None to indicate an error
  476. return [good_turns_time, cumulated_good_turns], [bad_turns_time, cumulated_bad_turns]
  477. def accumulation_threshold(cummulated_events_time, threshold_fraction=0.8):
  478. """
  479. Compute the first time point at which the cummulated number of event in cummulated_events_time
  480. is above a threshold, that is a fraction (80% by default) of the total number of events.
  481. Arguments:
  482. cummulated_events_time (list): list of time at which one event occurs.
  483. threshold_fraction (float, optional): fraction of the total number of event that will used as threshold.
  484. Returns:
  485. float: first value in cummulated_events_time that is above the
  486. """
  487. # Compute the total number of events
  488. total_events = len(cummulated_events_time)
  489. # Initialise index of the first event
  490. i = 0
  491. # Compute the fraction of cummulated events
  492. fraction_of_events = (i+1)/total_events if total_events!=0 else 0
  493. # Iterate while the fraction of cummulated events is lower than the threshold fraction
  494. while fraction_of_events<threshold_fraction and total_events!=0:
  495. # Increase the index
  496. i = i + 1
  497. # Update the fraction of cummulated events at the new index
  498. fraction_of_events = (i+1)/total_events
  499. res = cummulated_events_time[i] if total_events!=0 else 0
  500. # Return the first time in cummulated_events_time at which the fraction of cummulated events is higher than threshold_fraction.
  501. return res
  502. # %% [markdown]
  503. # ### Here will be computed and stored metrics related to the number of QTs that will be plotted for each mouse across sessions. Those metrics are:
  504. # - The number of turns in rewarding direction
  505. # - The ratio of the difference of the number of CCW and CW turns with respect to the total number of turns. We call it the CCW vs CW normalized difference.
  506. # $ \frac{N_{CCW}-N_{CW}}{N_{CCW}+N_{CW}} $
  507. # - The time that the mouse took to perform 80% of the total number of turns in rewarded direction of the session.
  508. # %%
  509. # Initialize dictionaries to store the various metrics for each mouse
  510. mice_rewarded_qts_persession = {mouse: [] for mouse in mice_to_analyse}
  511. mice_unrewarded_qts_persession = {mouse: [] for mouse in mice_to_analyse}
  512. mice_ccw_vs_cw_norm_diff_persession = {mouse: [] for mouse in mice_to_analyse}
  513. mice_qts_rewarded_dir_threshold_persession = {mouse: [] for mouse in mice_to_analyse}
  514. # Iterate through each mouse to process its data
  515. for mouse in mice_to_analyse:
  516. folder_path_mouse_to_process = os.path.join(path_to_data_folder, mouse)
  517. # Get the list of sessions for the current mouse
  518. sessions_to_process = sorted([name for name in os.listdir(folder_path_mouse_to_process)
  519. if os.path.isdir(os.path.join(folder_path_mouse_to_process, name))
  520. and name.startswith('MOU')])
  521. # limit the analysis to the subseet of session we want to analyse
  522. sessions_to_process = sessions_to_process[first_and_last_session_indexes[0]:first_and_last_session_indexes[1]]
  523. nb_sessions = len(sessions_to_process)
  524. print(f'Mouse {mouse}. There is/are {nb_sessions} sessions:')
  525. print(sessions_to_process, '\n')
  526. # Process each session for the current mouse
  527. for session_index, session_to_process in enumerate(sessions_to_process):
  528. # Define the pickle file path for the session
  529. output_pickle_filename = f"{session_to_process}_basic_processing_output.pickle"
  530. output_pickle_filepath = os.path.join(folder_path_mouse_to_process, session_to_process, output_pickle_filename)
  531. # Check if the pickle file exists
  532. if not os.path.exists(output_pickle_filepath):
  533. print(f'Pickle file does not exist for session {session_to_process}, skipping .....')
  534. # Append session Nan data to the respective dictionaries
  535. mice_ccw_vs_cw_norm_diff_persession[mouse].append([session_index + 1, np.nan])
  536. mice_rewarded_qts_persession[mouse].append([session_index + 1, np.nan])
  537. mice_unrewarded_qts_persession[mouse].append([session_index + 1, np.nan])
  538. mice_qts_rewarded_dir_threshold_persession[mouse].append([session_index + 1, np.nan])
  539. continue # Skip to the next session if the pickle file does not exist
  540. # Load the pickle file
  541. with open(output_pickle_filepath, 'rb') as file:
  542. session_data = pickle.load(file)
  543. # Determine the rewarded direction for all sessions of the current mouse
  544. rewarded_direction = finding_mouse_rewarded_direction(folder_path_mouse_to_process, first_and_last_session_indexes[0]+session_index)
  545. # Extract run around tower results from the session data
  546. epochs = session_data['all_epochs']
  547. qts = filter_qts(epochs['run_around_tower'])
  548. total_qts, cw_qts, rewarded_qts = count_qts(qts)
  549. ccw_qts = total_qts - cw_qts
  550. ccw_vs_cw_norm_diff = (ccw_qts-cw_qts)/(total_qts) if total_qts!=0 else 0
  551. # Compute cumulated turns time profiles for turns in rewarded and unrewarded direction, if there is a unique rewarded direction
  552. [rewarded_dir_turns_time, cumulated_rewarded_dir_turns], [unrewarded_dir_turns_time, cumulated_unrewarded_dir_turns] = cumulated_turns_time_profile(folder_path_mouse_to_process, session_to_process, rewarded_direction)
  553. # Determine the rewarded direction turns threshold time
  554. rewarded_dir_turns_threshold_time = accumulation_threshold(rewarded_dir_turns_time)
  555. # If reward delivery is not allowed, set to numpy.nan the metrics values concerning rewarded direction
  556. if rewarded_direction == 'X':
  557. # Append session data to the respective dictionaries
  558. mice_ccw_vs_cw_norm_diff_persession[mouse].append([session_index + 1, ccw_vs_cw_norm_diff])
  559. mice_rewarded_qts_persession[mouse].append([session_index + 1, np.nan])
  560. mice_unrewarded_qts_persession[mouse].append([session_index + 1, np.nan])
  561. mice_qts_rewarded_dir_threshold_persession[mouse].append([session_index + 1, np.nan])
  562. else:
  563. # Append session data to the respective dictionaries
  564. mice_ccw_vs_cw_norm_diff_persession[mouse].append([session_index + 1, ccw_vs_cw_norm_diff])
  565. mice_rewarded_qts_persession[mouse].append([session_index + 1, rewarded_qts])
  566. mice_unrewarded_qts_persession[mouse].append([session_index + 1, total_qts-rewarded_qts])
  567. mice_qts_rewarded_dir_threshold_persession[mouse].append([session_index + 1, rewarded_dir_turns_threshold_time])
  568. # %% [markdown]
  569. # # 6. Computations for panel F
  570. # %% [markdown]
  571. # ### The following cells will be used to compute the median of the median Fréchet distance of each pair of flattened trajectories;
  572. # %%
  573. def towers_coordinates_as_dictionnary(towers_coordinates_cm):
  574. """
  575. Converts a dictionary of tower coordinates into a structured dictionary
  576. where each tower's coordinates are labeled with explicit corner names.
  577. Arguments:
  578. towers_coordinates_cm (dict): Dictionary where keys are tower names and
  579. values are lists/tuples of four coordinates.
  580. Returns:
  581. dict: A dictionary mapping tower names to their respective coordinates,
  582. labeled as 'NW' (North-West), 'NE' (North-East), 'SE' (South-East),
  583. and 'SW' (South-West).
  584. """
  585. # Initialize a dictionary to store labeled coordinates
  586. towers_coordinates_as_dict = {}
  587. # Predefined corner names in the order expected from input coordinates
  588. corner_names = ['NW', 'NE', 'SE', 'SW']
  589. # Map each tower's coordinates to its corresponding corner names
  590. for tower, coordinates in towers_coordinates_cm.items():
  591. towers_coordinates_as_dict[tower] = {
  592. corner_names[i]: coord for i, coord in enumerate(coordinates)
  593. }
  594. return towers_coordinates_as_dict
  595. def get_tower_and_corner(run_around_tower):
  596. """
  597. Get the label of the tower and corner around wich the run around tower is happened
  598. based on the second and third elements saved in run_around_tower: 'N' for north, 'S' for south, 'E' for east, 'W' for west.
  599. Argument:
  600. run_around_tower (list): a list containing information about a run around tower
  601. Returns:
  602. str: Name of the tower around which the run occured
  603. str: A two character string. The first character is the name of the starting trapeze,
  604. the last character is the name of the ending trapeze
  605. """
  606. # Extract tower name, starting name and ending trapeze name
  607. tower_name = run_around_tower[1][0]
  608. start_trapeze = run_around_tower[1][1]
  609. end_trapeze = run_around_tower[2][1]
  610. # Determine the corner based on the start and end faces
  611. if start_trapeze == 'W' and end_trapeze == 'S':
  612. corner = 'SW'
  613. elif start_trapeze == 'S' and end_trapeze == 'E':
  614. corner = 'SE'
  615. elif start_trapeze == 'E' and end_trapeze == 'N':
  616. corner = 'NE'
  617. elif start_trapeze == 'N' and end_trapeze == 'W':
  618. corner = 'NW'
  619. elif start_trapeze == 'W' and end_trapeze == 'N':
  620. corner = 'NW'
  621. elif start_trapeze == 'N' and end_trapeze == 'E':
  622. corner = 'NE'
  623. elif start_trapeze == 'E' and end_trapeze == 'S':
  624. corner = 'SE'
  625. elif start_trapeze == 'S' and end_trapeze == 'W':
  626. corner = 'SW'
  627. else:
  628. corner = None # Handle unexpected cases
  629. return tower_name, corner
  630. def rotate_sw_trajectory_90_ccw(trajectory):
  631. """
  632. Rotates the input trajectory of 90° in counter-clockwise direction.
  633. Argument:
  634. trajectory (numpy.ndarray): A 2D numpy.ndarray of shape (2, X).
  635. Returns:
  636. numpy.ndarray: trajectory rotated by 90° counter-clockwise
  637. """
  638. # Define the roation matrix
  639. rotation_matrix = np.array([[0, -1], [1, 0]])
  640. # Returns the matrix product of the rotation matrix and the trajectory
  641. return rotation_matrix @ trajectory
  642. def rotate_nw_trajectory_180_ccw(trajectory):
  643. """
  644. Rotates the input trajectory of 180° in counter-clockwise direction.
  645. Argument:
  646. trajectory (numpy.ndarray): A 2D numpy.ndarray of shape (2, X).
  647. Returns:
  648. numpy.ndarray: trajectory rotated by 180° counter-clockwise
  649. """
  650. return -trajectory
  651. def rotate_ne_trajectory_270_ccw(trajectory):
  652. """
  653. Rotates the input trajectory of 270° in counter-clockwise direction.
  654. Argument:
  655. trajectory (numpy.ndarray): A 2D numpy.ndarray of shape (2, X).
  656. Returns:
  657. numpy.ndarray: trajectory rotated by 270° counter-clockwise
  658. """
  659. # Define the roation matrix
  660. rotation_matrix = np.array([[0, 1], [-1, 0]])
  661. # Returns the matrix product of the rotation matrix and the trajectory
  662. return rotation_matrix @ trajectory
  663. def compute_frechet_distances(runs_list):
  664. """
  665. This function takes a list of trajectories of dimensions (2,N), computes the Fréchet distance
  666. between each pairs of trajectories, and returns the median of those distances.
  667. Arguments:
  668. runs_list (list): List of trajectories of dimensions (2,N)
  669. Return:
  670. (float): Median of the Fréchet distances of all trajectories pairs
  671. """
  672. # Create an empty list that will contain the Fréchet distances of all trajectories pairs
  673. frechet_distances_list = []
  674. # Iterate on all trajectories
  675. for i in range(len(runs_list)):
  676. # Transpose the trajectory
  677. # This step is necessary as the function that computes the Fréchet distance takes as an argument trajectories of dimensions (N,2)
  678. traj_points_a = np.transpose(runs_list[i])
  679. # Iterate on every trajectories of higher ranks to create trajectories pairs
  680. # We do that so that no pairs is counted twice, as the Fréchet distance between trajectory A and B is the same as between B and A
  681. for j in range(i+1, len(runs_list)):
  682. # Transpose the trajectory
  683. traj_points_b = np.transpose(runs_list[j])
  684. # Compute the Fréchet distance for the current pair of trajectories
  685. frechet_distance = sm.frechet_dist(traj_points_a,traj_points_b)
  686. # Store the result in a list
  687. frechet_distances_list.append(frechet_distance)
  688. # Return the median of the list of Fréchet distances
  689. return frechet_distances_list
  690. # %% [markdown]
  691. # ### This cell process the runs around tower trajectory such that:
  692. # - All runs around tower with the same starting and ending trapeze and are shifted to have their origin on the same axis, independently of the tower where they occur.
  693. # - All runs around tower are rotated from 90°/180°/270° counter-clockwise when the corner around which the mouse turns is South-West/North-West/North-East.
  694. # - Does so separately for turns in clockwise and counter-clockwise direction.
  695. # %%
  696. # Initialize a dictionary to store all trajectories for each mouse, with empty lists for each session
  697. mice_alltrajectories_persession = {mouse: {} for mouse in mice_to_analyse}
  698. # Initialize a dictionary to store realigned and rotated trajectories for each mouse, with empty dictionaries for each session
  699. trajectories_per_session_realigned_rotated = {mouse: {} for mouse in mice_to_analyse}
  700. # Loop through each mouse in the list of mice to realign all their turns trajectory
  701. for mouse in mice_to_analyse:
  702. # Define the folder path for the current mouse
  703. folder_path_mouse_to_process = os.path.join(path_to_data_folder, mouse)
  704. # Get a sorted list of session folders for the current mouse
  705. sessions_to_process = sorted([name for name in os.listdir(folder_path_mouse_to_process)
  706. if os.path.isdir(os.path.join(folder_path_mouse_to_process, name))
  707. and name.startswith('MOU')])
  708. sessions_to_process = sessions_to_process[first_and_last_session_indexes[0]:first_and_last_session_indexes[1]]
  709. # Get the number of sessions to process
  710. nb_sessions = len(sessions_to_process)
  711. print(f'Processing mouse {mouse}. There is/are {nb_sessions} sessions to process:')
  712. print(sessions_to_process, '\n')
  713. # Loop through each session for the current mouse to realign all the turns trajectory from those sessions
  714. for session_index, session_to_process in enumerate(sessions_to_process):
  715. print(f'Getting the run trajectory of session {session_index}')
  716. # Define the path to the pickle file for the current session
  717. output_pickle_filename = f"{session_to_process}_basic_processing_output.pickle"
  718. output_pickle_filepath = os.path.join(folder_path_mouse_to_process, session_to_process, output_pickle_filename)
  719. # Check if the pickle file exists
  720. if not os.path.exists(output_pickle_filepath):
  721. mice_alltrajectories_persession[mouse][session_index] = []
  722. trajectories_per_session_realigned_rotated[mouse][session_index] = {'CW':[], 'CCW':[]}
  723. print(f'Pickle file does not exist for session {session_to_process}, skipping .....')
  724. continue
  725. # Load the data from the pickle file
  726. with open(output_pickle_filepath, 'rb') as file:
  727. session_data = pickle.load(file)
  728. # Initialize entries for the current session in the dictionaries
  729. mice_alltrajectories_persession[mouse][session_index] = []
  730. trajectories_per_session_realigned_rotated[mouse][session_index] = {'CW':[], 'CCW':[]}
  731. # Get the runs around the tower for the current session
  732. runs_around_tower = filter_qts(session_data['all_epochs']['run_around_tower'])
  733. # Get the trajectory of the mouse
  734. positions = np.array(session_data['positions'])
  735. # Get the tower coordinates and convert them to a dictionary format
  736. towers_coordinates_cm = session_data['towers_coordinates_cm']
  737. towers_coordinates_as_dict = towers_coordinates_as_dictionnary(towers_coordinates_cm)
  738. # Loop through each run around the tower to realign them
  739. for run_around_tower in runs_around_tower:
  740. # Only process runs around tower where there was only one trapeze switch
  741. if run_around_tower[3]['num_trapezeswitch'] == 1:
  742. # Extract the run trajectory
  743. run_trajectory = positions[:, run_around_tower[0][0]:run_around_tower[0][1]]
  744. mice_alltrajectories_persession[mouse][session_index].append(run_trajectory)
  745. # Get the tower and corner names
  746. tower_name, corner = get_tower_and_corner(run_around_tower)
  747. # Access tower and corner coordinates from the dictionary
  748. if tower_name in towers_coordinates_as_dict and corner in towers_coordinates_as_dict[tower_name]:
  749. this_corner_coordinates = towers_coordinates_as_dict[tower_name][corner]
  750. else:
  751. print(f"Invalid tower or corner: {tower_name}, {corner}")
  752. continue
  753. # Extract the trajectory slice based on the start and end time indices
  754. start_idx, end_idx = run_around_tower[0]
  755. this_trajectory = positions[:, start_idx:end_idx]
  756. # Get the corner's reference coordinates (X and Y)
  757. newXreference = this_corner_coordinates[0]
  758. newYreference = this_corner_coordinates[1]
  759. # Shift the trajectory to reference the new corner coordinates
  760. this_trajectory[0, :] -= newXreference # Shift X coordinates
  761. this_trajectory[1, :] -= newYreference # Shift Y coordinates
  762. # Rotate the trajectory based on the corner
  763. if corner == 'SW':
  764. this_trajectory = rotate_sw_trajectory_90_ccw(this_trajectory)
  765. elif corner == 'NW':
  766. this_trajectory = rotate_nw_trajectory_180_ccw(this_trajectory)
  767. elif corner == 'NE':
  768. this_trajectory = rotate_ne_trajectory_270_ccw(this_trajectory)
  769. # Get the direction (CW or CCW)
  770. direction = run_around_tower[3]['direction']
  771. # Append the transformed trajectory to the appropriate list based on direction
  772. trajectories_per_session_realigned_rotated[mouse][session_index][direction].append(this_trajectory)
  773. # %% [markdown]
  774. # ### This cell uses realigned trajectories to compute the median of their Fréchet distances. It does so in several steps:
  775. # 1. Select a direction for the and loop on all mice
  776. # 2. Computes the Fréchet distance of all the pairs of trajectory in the given direction
  777. # 3. Computes the median of those coefficients and stores it in the corresponding dictionnary, depending on whether it's a turn in the rewarded or unrewarded direction.
  778. # 4. Does the same for the other direction.
  779. # %%
  780. #modified by david to try to save trajecotry Fréchet distance
  781. overwritte_previous_frechet_distances=False
  782. # Initialize dictionaries to store overall Fréchet distances per session for each direction (CW and CCW)
  783. overall_trajectory_frechet_distances_per_session = {mouse: {'CW': [], 'CCW': []} for mouse in trajectories_per_session_realigned_rotated}
  784. overall_cw_turns_frechet_distances_per_session = {mouse: [] for mouse in trajectories_per_session_realigned_rotated}
  785. overall_ccw_turns_frechet_distances_per_session = {mouse: [] for mouse in trajectories_per_session_realigned_rotated}
  786. allsessionnumber = list(range(first_and_last_session_indexes[0]+1, first_and_last_session_indexes[1]+1))
  787. # Define the directions to process
  788. directions = ['CW', 'CCW']
  789. # Loop through each mouse in the list of mice to analyze
  790. for mouse in mice_to_analyse:
  791. # Define the folder path for the current mouse
  792. folder_path_mouse_to_process = os.path.join(path_to_data_folder, mouse)
  793. # Get a sorted list of session folders for the current mouse
  794. sessions_to_process = sorted([name for name in os.listdir(folder_path_mouse_to_process)
  795. if os.path.isdir(os.path.join(folder_path_mouse_to_process, name))
  796. and name.startswith('MOU')])
  797. sessions_to_process = sessions_to_process[first_and_last_session_indexes[0]:first_and_last_session_indexes[1]]
  798. # Get the number of sessions to process for the current mouse
  799. nb_sessions = len(sessions_to_process)
  800. print(f'Processing mouse {mouse}. There is/are {nb_sessions} sessions to process:')
  801. print(sessions_to_process, '\n')
  802. # Loop through each session for the current mouse
  803. for session_index in trajectories_per_session_realigned_rotated[mouse]:
  804. # Define the pickle filename to save the realigned and rotated trajectories for this session
  805. session_to_process = sessions_to_process[session_index]
  806. overall_trajectory_frechet_distances_per_session_pickle_filename = (
  807. f"{session_to_process}_overall_trajectory_frechet_distances_per_session.pickle"
  808. )
  809. overall_trajectory_frechet_distances_per_session_pickle_filepath = os.path.join(
  810. folder_path_mouse_to_process, session_to_process, overall_trajectory_frechet_distances_per_session_pickle_filename
  811. )
  812. if not os.path.exists(overall_trajectory_frechet_distances_per_session_pickle_filepath) or overwritte_previous_frechet_distances:
  813. this_session_frechet_distance = {'CW': None, 'CCW': None}
  814. print(f"Processing session index: {session_index} for direction: {direction}")
  815. # Loop through each direction (CW and CCW)
  816. for direction in directions:
  817. # Access the realigned trajectories for the current session and direction
  818. realigned_trajectories = trajectories_per_session_realigned_rotated[mouse][session_index][direction]
  819. # If there are no trajectories for the current direction, skip the session
  820. if not realigned_trajectories:
  821. print(f"No trajectories for {direction} in session {session_index}, skipping...")
  822. overall_trajectory_frechet_distances_per_session[mouse][direction].append([allsessionnumber[session_index], np.nan])
  823. this_session_frechet_distance[direction]=np.nan
  824. continue
  825. # Compute the number of realigned trajectories
  826. number_of_realigned_trajectories = len(realigned_trajectories)
  827. # If there are fewer than 5 trajectories, skip the session
  828. if number_of_realigned_trajectories < 5:
  829. print(f"Less than 5 trajectories for {direction} in session {session_index}, skipping...")
  830. overall_trajectory_frechet_distances_per_session[mouse][direction].append([allsessionnumber[session_index], np.nan])
  831. if direction=='CW':
  832. this_session_frechet_distance['CW']=np.nan
  833. if direction=='CCW':
  834. this_session_frechet_distance['CCW']=np.nan
  835. continue
  836. else:
  837. # Compute Fréchet distance
  838. frechet_distances = compute_frechet_distances(trajectories_per_session_realigned_rotated[mouse][session_index][direction])
  839. # Compute the overall Fréchet distance (median of pairwise frechet_distances)
  840. overall_frechet_distance = np.nanmedian(frechet_distances)
  841. # Append the overall trajectory Fréchet distance for the current direction and session
  842. overall_trajectory_frechet_distances_per_session[mouse][direction].append([allsessionnumber[session_index], overall_frechet_distance])
  843. if direction=='CW':
  844. this_session_frechet_distance['CW']=overall_frechet_distance
  845. if direction=='CCW':
  846. this_session_frechet_distance['CCW']=overall_frechet_distance
  847. # Save the dictionary for CW and CCW trajectories
  848. with open(overall_trajectory_frechet_distances_per_session_pickle_filepath, 'wb') as file:
  849. pickle.dump(this_session_frechet_distance, file)
  850. print(f"Saved trajectories Fréchet distances for {mouse}, session {session_to_process} at {overall_trajectory_frechet_distances_per_session_pickle_filepath}")
  851. else:
  852. with open(overall_trajectory_frechet_distances_per_session_pickle_filepath, 'rb') as file:
  853. this_session_frechet_distance = pickle.load(file) # Preserve existing values
  854. print(f"Loaded existing Fréchet distances for {mouse}, session {session_to_process}")
  855. for direction in directions:
  856. overall_frechet_distance=this_session_frechet_distance[direction]
  857. overall_trajectory_frechet_distances_per_session[mouse][direction].append([allsessionnumber[session_index], overall_frechet_distance])
  858. # Store the all the Fréchet distances and session in the corresponding direction's dictionnary
  859. for direction in directions:
  860. if direction=='CW':
  861. overall_cw_turns_frechet_distances_per_session[mouse] = overall_trajectory_frechet_distances_per_session[mouse][direction]
  862. if direction=='CCW':
  863. overall_ccw_turns_frechet_distances_per_session[mouse] = overall_trajectory_frechet_distances_per_session[mouse][direction]
  864. # %% [markdown]
  865. # # 7. Computations for panel H
  866. # %% [markdown]
  867. # ### This cell computes and store the median of the turns maximum speed.
  868. # %%
  869. # Initialize dictionaries to store the various metrics for each mouse
  870. mice_median_maximum_cw_turn_speed_persession = {mouse: [] for mouse in mice_to_analyse}
  871. mice_median_maximum_ccw_turn_speed_persession = {mouse: [] for mouse in mice_to_analyse}
  872. # Iterate through each mouse to process its data
  873. for mouse in mice_to_analyse:
  874. folder_path_mouse_to_process = os.path.join(path_to_data_folder, mouse)
  875. # Get the list of sessions for the current mouse
  876. sessions_to_process = sorted([name for name in os.listdir(folder_path_mouse_to_process)
  877. if os.path.isdir(os.path.join(folder_path_mouse_to_process, name))
  878. and name.startswith('MOU')])
  879. # limit the analysis to the subseet of session we want to analyse
  880. sessions_to_process = sessions_to_process[first_and_last_session_indexes[0]:first_and_last_session_indexes[1]]
  881. nb_sessions = len(sessions_to_process)
  882. print(f'Mouse {mouse}. There is/are {nb_sessions} sessions:')
  883. print(sessions_to_process, '\n')
  884. # Process each session for the current mouse
  885. for session_index, session_to_process in enumerate(sessions_to_process):
  886. # Define the pickle file path for the session
  887. output_pickle_filename = f"{session_to_process}_basic_processing_output.pickle"
  888. output_pickle_filepath = os.path.join(folder_path_mouse_to_process, session_to_process, output_pickle_filename)
  889. # Check if the pickle file exists
  890. if not os.path.exists(output_pickle_filepath):
  891. print(f'Pickle file does not exist for session {session_to_process}, skipping .....')
  892. mice_median_maximum_cw_turn_speed_persession[mouse].append([session_index + 1, np.nan])
  893. mice_median_maximum_ccw_turn_speed_persession[mouse].append([session_index + 1, np.nan])
  894. continue # Skip to the next session if the pickle file does not exist
  895. # Load the pickle file
  896. with open(output_pickle_filepath, 'rb') as file:
  897. session_data = pickle.load(file)
  898. # Extract run around tower results from the session data
  899. run_around_tower_sessionresult = session_data['run_around_tower_sessionresult']
  900. # Initialize lists to store turn speeds
  901. cw_turns_max_speed = []
  902. ccw_turns_max_speed = []
  903. runs = filter_qts(session_data["all_epochs"]["run_around_tower"])
  904. # Iterate through each run around the tower
  905. for run in runs:
  906. # Append turn speeds in the corresponding direction's dictionnary
  907. if run[3]['direction'] == 'CW':
  908. cw_turns_max_speed.append(run[4]["epoch_maxspeed"])
  909. if run[3]['direction'] == 'CCW':
  910. ccw_turns_max_speed.append(run[4]["epoch_maxspeed"])
  911. # Append session data to the respective dictionaries
  912. mice_median_maximum_cw_turn_speed_persession[mouse].append([session_index + 1, np.median(cw_turns_max_speed)])
  913. mice_median_maximum_ccw_turn_speed_persession[mouse].append([session_index + 1, np.median(ccw_turns_max_speed)])
  914. # %% [markdown]
  915. # # 8. Plot Panel A
  916. # %% [markdown]
  917. # The function detect_clean_run_epochs exists in the rocessing_session_trajectory.py script, but only returns the clean epochs (i.e taking into account the initiations of runs, not only the time when the mouse is going above a speed threshold). The function detect_raw_and_clean_run_epochs has been defined to return both clean and raw epochs (i.e the epochs defined only with a speed threshold).
  918. # %%
  919. def detect_raw_and_clean_run_epochs(speeds, time):
  920. """
  921. Identifies continuous epochs during which the mouse is moving above a certain speed (cut_off_speed).
  922. A minimal duration of low speed is necessary to be considered as the end of a run.
  923. Similarly, a minimal duration of high speed is necessary to be considered as a run.
  924. Arguments:
  925. speeds (list): speed of the mouse at every point of the trajectory
  926. time_video_frames (list): time of every point of the trajectory
  927. Returns:
  928. run_epochs (list): List of frame intervals during which a run occurs accordin to the speed threshold
  929. clean_run_epochs (list): List of frame intervals during which a run occurs, extended with frames adjacent to the first/last one
  930. where the acceleration is at least 10% of the acceleration when going above/below the speed threshold.
  931. """
  932. # For this we need some parameters to cut the trajectory into run based on speed, duration of runs and pauses
  933. pause_min_duration = 0.1 # If a stop is shorter than this, merges the two epochs bordering it
  934. run_min_duration = 0.3 # Minimal duration of an epoch to be considerd
  935. cut_off_speed = 7 # This value is the speed in cm/s. It is used to detect when the animals stop running.
  936. # Create a list to store run epochs
  937. run_epochs = []
  938. # Flag to track if we are currently in a running epoch
  939. is_in_epoch = False
  940. # Initialise the starting index of the first epoch
  941. epoch_start_index = 0
  942. # Raises an error if the speed and time vectors don't have the same number of values
  943. if len(speeds) != len(time):
  944. raise ValueError("speeds and time_video_frames have different lengths")
  945. # Iterate on speed
  946. for i in range(len(speeds)):
  947. # Check if speed above cut-off value
  948. if speeds[i] >= cut_off_speed:
  949. # Check if the previous trajectory speed was part of running epoch. If not this will be a start of a new epoch
  950. if not is_in_epoch:
  951. # Mark the beginning of a new epoch
  952. epoch_start_index = i
  953. # Change flag to indicate the current speed belongs to a the current epoch
  954. is_in_epoch = True
  955. else: # the speed of the current data point is below the treshold
  956. if is_in_epoch: # if we were in a run epoch just before (1st point below the treshold)
  957. # Check first if the pause between this epoch's starting point (time_video_frames[epoch_start_index]) and
  958. # the previous epoch's last point time_video_frames[run_epochs[-1][1]] is shorter than the minimal time for a pause.
  959. # If it is, the previous epoch should be extended to the previous data point.
  960. if run_epochs and (time[epoch_start_index] - time[run_epochs[-1][1]] < pause_min_duration):
  961. run_epochs[-1][1] = i - 1 # Extend the previous epoch
  962. else: # the pause has been long enough then we terminate the run epoch other previous
  963. run_epochs.append([epoch_start_index, i - 1]) # Add new epoch
  964. # Change flag to indicate the current speed belongs to a new epoch
  965. is_in_epoch = False
  966. # Check for any epoch still in progress after having processed the last speed value
  967. if is_in_epoch:
  968. # Check if run_epochs is not empty and the difference between the last time of the trajectory and the epoch start time is lower than pause_min_duration
  969. if run_epochs and (time[epoch_start_index] - time[run_epochs[-1][1]] < pause_min_duration):
  970. # If it is not, sets the last speed value as its end
  971. run_epochs[-1][1] = len(speeds) - 1
  972. # Else, check if the difference between the last time of the trajectory and the epoch start time is higher than pause_min_duration
  973. elif (time[-1] - time[epoch_start_index]) >= run_min_duration:
  974. run_epochs.append([epoch_start_index, len(speeds) - 1])
  975. # Remove epochs that are too short
  976. run_epochs = [epoch for epoch in run_epochs if (time[epoch[1]] - time[epoch[0]]) >= run_min_duration]
  977. # Adjust the start and end of each epoch based on acceleration, to take into account the initiation of the movement.
  978. # Find the point at wich the animal acceleration is less than 10% of the acceleration at the moment
  979. # at which it goes above the speed treshold.
  980. # Initialize a list of None of the length of run_epochs, that will contain the run epochs extended by the method described above.
  981. clean_run_epochs = [None] * len(run_epochs)
  982. # Iterate on all epochs
  983. for index,epoch in enumerate(run_epochs):
  984. # Copy the current epoch in clean_run_epochs at the same index
  985. clean_run_epochs[index] = epoch.copy()
  986. # Define the start and end index of the epoch
  987. epoch_start, epoch_end = epoch[0], epoch[1]
  988. # Compute the acceleration when going above the speed threshold
  989. current_point = epoch_start
  990. acceleration_at_crossing=(speeds[current_point + 1] - speeds[current_point]) / (time[current_point + 1] - time[current_point])
  991. # Iterate on indexes backward from the start of the epoch
  992. while current_point > 0:
  993. # Compute the acceleration at the previous index
  994. previous_acceleration = (speeds[current_point] - speeds[current_point - 1]) / (time[current_point] - time[current_point - 1])
  995. # Check if this acceleration is lower than or equal to 10% of the crossing acceleration
  996. if previous_acceleration <= (0.1 * acceleration_at_crossing) or previous_acceleration <= 0:
  997. # If it is, stops to iterate on indexes. current_point is then the last index before acceleration is higher than or equal to 10%
  998. break
  999. current_point -= 1
  1000. # Set the beginning of the epoch at current_point
  1001. clean_run_epochs[index][0] = current_point
  1002. # Adjust the end of the epoch
  1003. # Find the point at wich the animal acceleration is less than 10% of the acceleration at the moment
  1004. # at which it goes above the speed treshold.
  1005. # Compute the acceleration when going bellow the speed threshold
  1006. current_point = epoch_end
  1007. acceleration_at_crossing=(speeds[current_point - 1] - speeds[current_point]) / (time[current_point] - time[current_point-1])
  1008. # Iterate on indexes forward from the end of the epoch
  1009. while current_point < len(speeds) - 1:
  1010. # Compute the acceleration at the next index
  1011. next_acceleration = (speeds[current_point] - speeds[current_point + 1]) / (time[current_point+1] - time[current_point])
  1012. # Check if this acceleration is lower than or equal to 10% of the crossing acceleration
  1013. if next_acceleration <= (0.1 * acceleration_at_crossing) or next_acceleration <= 0:
  1014. # If it is, stops to iterate on indexes. current_point is then the last point before acceleration is lower than or equal to 10%
  1015. break
  1016. current_point += 1
  1017. # Set the beginning of the epoch at current_point
  1018. clean_run_epochs[index][1] = current_point
  1019. return run_epochs, clean_run_epochs
  1020. def find_run_type(run_epoch, folder_path_mouse_to_analyse, session_index, start_session_index=0):
  1021. """
  1022. Determines the type of run epoch by loading session data from a pickle file.
  1023. Arguments:
  1024. run_epoch (int): The epoch ID to classify.
  1025. folder_path_mouse_to_analyse (str): Path to the folder containing session data.
  1026. session_index (int): Index of the session from which the run epoch will be classified.
  1027. start_session_index (int, optional): index of the first session in the list of sessions to analyse.
  1028. If different from 0, the session index should refer to
  1029. the index of the session in the sub-list of session to process, not the total list
  1030. Returns:
  1031. epoch_type (str): The classified type of the run epoch (e.g., 'run_around_tower', 'exploratory_run', etc.).
  1032. """
  1033. # Get all session folders that start with 'MOU' and sort them
  1034. sessions_to_analyse = sorted([name for name in os.listdir(folder_path_mouse_to_analyse)
  1035. if os.path.isdir(os.path.join(folder_path_mouse_to_analyse, name))
  1036. and name.startswith('MOU')])
  1037. session_to_analyse = sessions_to_analyse[start_session_index+session_index]
  1038. # Define the pickle file path
  1039. output_pickle_filename = f"{session_to_analyse}_basic_processing_output.pickle"
  1040. output_pickle_filepath = os.path.join(folder_path_mouse_to_analyse, session_to_analyse, output_pickle_filename)
  1041. # Load the pickle file containing session data
  1042. with open(output_pickle_filepath, 'rb') as file:
  1043. session_data = pickle.load(file)
  1044. # Extract epoch start and end frames, and consider it is their IDs for different movement types
  1045. run_around_tower_ids = np.array([item[0] for item in filter_qts(session_data['all_epochs']['run_around_tower'])])
  1046. run_between_towers_ids = np.array([item[0] for item in session_data['all_epochs']['run_between_towers']])
  1047. exploratory_run_ids = np.array([item[0] for item in session_data['all_epochs']['exploratory_run']])
  1048. run_toward_tower_ids = np.array([item[0] for item in session_data['all_epochs']['run_toward_tower']])
  1049. # Extract immobility epochs (which start and )
  1050. immobility_ids = np.array([[item[0], item[1]] for item in session_data['all_epochs']['immobility']])
  1051. # Check if the given run_epoch belongs to any of the predefined categories
  1052. condition_1 = run_epoch in run_around_tower_ids
  1053. condition_2 = run_epoch in run_between_towers_ids
  1054. condition_3 = run_epoch in run_toward_tower_ids
  1055. condition_4 = run_epoch in exploratory_run_ids
  1056. condition_5 = run_epoch in immobility_ids
  1057. # Determine the type of the run epoch
  1058. if np.any(condition_1):
  1059. epoch_type = "run_around_tower"
  1060. elif np.any(condition_2):
  1061. epoch_type = "run_between_towers"
  1062. elif np.any(condition_3):
  1063. epoch_type = "run_toward_tower"
  1064. elif np.any(condition_4):
  1065. epoch_type = "exploratory_run"
  1066. elif np.any(condition_5):
  1067. epoch_type = "immobility"
  1068. else:
  1069. print("WARNING: unclassified epoch")
  1070. epoch_type = "unclassified"
  1071. return epoch_type
  1072. def plot_selected_run_epochs(folder_path_mouse_to_analyse, session_index, first_epoch_to_plot, last_epoch_to_plot, arena_coordinates, ax, start_session_index=0, points_for_direction=6, show_legend=True):
  1073. """
  1074. Plots selected run epochs from a session.
  1075. Arguments:
  1076. folder_path_mouse_to_analyse (str): Path to the folder containing session data.
  1077. session_index (int): Index of the session to analyse.
  1078. first_epoch_to_plot (int): Index of the first epoch to plot.
  1079. last_epoch_to_plot (int): Index of the last epoch to plot.
  1080. arena_coordinates (list): List of four 2D coordinates, locating the corners of the arena.
  1081. ax (matplotlib.axes.Axes): Matplotlib axis object for plotting.
  1082. start_session_index (int, optional): index of the first session in the list of sessions to analyse.
  1083. If different from 0, the session index should refer to
  1084. the index of the session in the sub-list of session to process, not the total list.
  1085. points_for_direction (int, optional): Number of points used to determine movement direction. Default is 4.
  1086. show_legend (bool, optional): Whether to show a legend. Default is True.
  1087. Returns:
  1088. (int): Start index of the first plotted epoch.
  1089. (int): End index of the last plotted epoch.
  1090. (list): List of trapeze switch times during the plotted epochs.
  1091. """
  1092. # Get all session folders that start with 'MOU' and sort them
  1093. sessions_to_analyse = sorted([name for name in os.listdir(folder_path_mouse_to_analyse)
  1094. if os.path.isdir(os.path.join(folder_path_mouse_to_analyse, name))
  1095. and name.startswith('MOU')])
  1096. session_to_analyse = sessions_to_analyse[start_session_index+session_index]
  1097. # Define the pickle file path
  1098. output_pickle_filename = f"{session_to_analyse}_basic_processing_output.pickle"
  1099. output_pickle_filepath = os.path.join(folder_path_mouse_to_analyse, session_to_analyse, output_pickle_filename)
  1100. # Load session data from the pickle file
  1101. with open(output_pickle_filepath, 'rb') as file:
  1102. session_data = pickle.load(file)
  1103. # Load additional data related to the session
  1104. traj_df, turns_df, param_df = load_data(folder_path_mouse_to_analyse, session_to_analyse)
  1105. # Extract session parameters
  1106. traject_time = session_data['timeofframes']
  1107. speeds = session_data['speeds']
  1108. smoothed_positions = session_data['positions']
  1109. smoothed_xpositions = smoothed_positions[0]
  1110. smoothed_ypositions = smoothed_positions[1]
  1111. all_trapezes_coordinates_cm = session_data['all_trapezes_coordinates_cm']
  1112. # Detect run epochs
  1113. clean_run_epochs = detect_run_epochs(speeds, traject_time)
  1114. #run_epochs, clean_run_epochs = detect_raw_and_clean_run_epochs(speeds, traject_time)
  1115. # Define run types and corresponding colors
  1116. run_types = ["run_around_tower", "run_between_towers", "run_toward_tower", "exploratory_run", "immobility", "unclassified"]
  1117. type_names = ["Quarter-Turns (QTs, CCW)", "Run Between Towers (RBTs)", "Run toward tower", "Exploratory run", "Immobility", "Unclassified"]
  1118. colors = ["lightcoral", "green", "orange", "blue", "yellow", "black"]
  1119. trapeze_switch_times = []
  1120. all_start_end_indexes = []
  1121. labels_displayed = set()
  1122. # Iterate through selected epochs
  1123. for idx in range(first_epoch_to_plot, last_epoch_to_plot):
  1124. run_epoch = clean_run_epochs[idx]
  1125. start_index, end_index = run_epoch[0], run_epoch[1]
  1126. all_start_end_indexes.extend([start_index, end_index])
  1127. # Check if indexes are within valid bounds
  1128. if start_index < 0 or end_index >= len(traject_time):
  1129. print(f"Indexes out of bounds for run_epoch: {run_epoch}")
  1130. continue
  1131. # Extract trajectory segment for the epoch
  1132. run_epoch = [smoothed_xpositions[start_index: end_index+1], smoothed_ypositions[start_index: end_index+1]]
  1133. times_run_epoch = traject_time[start_index: end_index+1]
  1134. # Determine epoch type
  1135. epoch_type = find_run_type([start_index, end_index], folder_path_mouse_to_analyse, session_index)
  1136. # Get corresponding color
  1137. color = colors[run_types.index(epoch_type)]
  1138. # Avoid duplicate legend labels
  1139. if color not in labels_displayed:
  1140. line_label = type_names[run_types.index(epoch_type)]
  1141. labels_displayed.add(color)
  1142. else:
  1143. line_label = ''
  1144. # Plot the trajectory segment
  1145. ax.plot(smoothed_xpositions[start_index: end_index+1], smoothed_ypositions[start_index: end_index+1], color=color, linestyle='-', label=line_label)
  1146. # Mark the beginning of the trajectory
  1147. ax.plot(run_epoch[0][0], run_epoch[1][0], marker='o', color='black', markersize=0.1)
  1148. # Draw an arrow indicating the direction of movement at the end of the last epoch
  1149. #if idx + 1 == last_epoch_to_plot:
  1150. dx = smoothed_xpositions[end_index+1] - smoothed_xpositions[end_index+1-points_for_direction]
  1151. dy = smoothed_ypositions[end_index+1] - smoothed_ypositions[end_index+1-points_for_direction]
  1152. norm_speed = np.hypot(dx, dy)
  1153. if norm_speed != 0:
  1154. dx /= norm_speed
  1155. dy /= norm_speed
  1156. ax.arrow(smoothed_xpositions[end_index+1], smoothed_ypositions[end_index+1], dx, dy,
  1157. head_width=2, head_length=2, fc='black', ec='black', zorder=100)
  1158. # Extract turns within this trajectory segment
  1159. turns_in_QT = turns_df[(turns_df['time'] >= times_run_epoch[0]) & (turns_df['time'] <= times_run_epoch[-1])]
  1160. trapeze_switch_times.extend(turns_in_QT['time'].values)
  1161. tower_coordinates = session_data['towers_coordinates_cm']
  1162. # Plot each tower
  1163. for i, (tower_name, vertices) in enumerate(tower_coordinates.items()):
  1164. tower_x, tower_y = zip(*vertices + [vertices[0]])
  1165. ax.fill(tower_x, tower_y, 'black', alpha=0.01)
  1166. ax.plot(tower_x, tower_y, 'k-', linewidth=0.5)
  1167. # Draw the arena perimeter
  1168. arena_x, arena_y = zip(*arena_coordinates + [arena_coordinates[0]])
  1169. ax.plot(arena_x, arena_y, 'grey', linewidth=1)
  1170. # Define fill colors for trapezes
  1171. fill_colors = ['lightgray'] * 4
  1172. # Plot trapezes
  1173. for i, (tower, trapezes) in enumerate(all_trapezes_coordinates_cm.items()):
  1174. for j, (trapeze, coordinates) in enumerate(trapezes.items()):
  1175. coordinates_copy = coordinates + [coordinates[0]] # Close the polygon
  1176. x_coords, y_coords = zip(*coordinates_copy)
  1177. ax.fill(x_coords, y_coords, color=fill_colors[j % len(fill_colors)], alpha=0.5)
  1178. # Remove axis spines and labels for a cleaner plot
  1179. for spine in ax.spines.values():
  1180. spine.set_visible(False)
  1181. ax.set_xticks([])
  1182. ax.set_yticks([])
  1183. # Display legend if required
  1184. if show_legend:
  1185. ax.legend(loc=[0, 1.3], frameon=False, ncols = 2)
  1186. return all_start_end_indexes[0], all_start_end_indexes[-1], trapeze_switch_times
  1187. def plot_trajectory_speed_chunk(start_idx, end_idx, traject_time, speeds, run_epochs, clean_run_epochs,
  1188. ax, folder_path_mouse_to_analyse, session_index, start_session_index=0, cut_off_speed=7
  1189. ):
  1190. """
  1191. Plots the trajectory speeds and highlights run epochs.
  1192. Arguments:
  1193. start_idx (int): Start index of the time window to plot.
  1194. end_idx (int): End index of the time window to plot.
  1195. traject_time (array-like): Array of time points corresponding to the trajectory.
  1196. speeds (array-like): Array of speed values at each time point.
  1197. run_epochs (list): List of tuples containing start and end indices of raw run epochs.
  1198. clean_run_epochs (list): List of tuples containing start and end indices of cleaned run epochs.
  1199. ax (matplotlib.axes.Axes): Matplotlib subplot axis to plot on.
  1200. folder_path_mouse_to_analyse (str): Path to the folder containing mouse session data.
  1201. session_index (int): Session identifier for the data.
  1202. cut_off_speed (float, optional): Speed threshold for highlighting certain speeds. Default is 7 cm/s.
  1203. """
  1204. # Plot the trajectory speeds with small markers and a thin line
  1205. ax.plot(
  1206. traject_time[start_idx:end_idx], speeds[start_idx:end_idx],
  1207. label='Trajectory Speeds', color='black', marker='o',
  1208. markerfacecolor='none', markersize=0.5, linewidth=0.5, zorder=11
  1209. )
  1210. # Draw a horizontal line indicating the cut-off speed threshold
  1211. ax.axhline(y=cut_off_speed, color='orange', linestyle='--', label='Cut-off Speed')
  1212. # Define run types and corresponding colors
  1213. run_types = ["run_around_tower", "run_between_towers", "run_toward_tower",
  1214. "exploratory_run", "immobility", "unclassified"]
  1215. colors = ["red", "green", "orange", "blue", "yellow", "black"]
  1216. # Highlight run epochs on the plot
  1217. for idx, clean_epoch in enumerate(clean_run_epochs):
  1218. epoch = run_epochs[idx] # Corresponding original (raw) run epoch
  1219. clean_epoch_start, clean_epoch_end = clean_epoch[0], clean_epoch[1]
  1220. epoch_start, epoch_end = epoch[0], epoch[1]
  1221. # Ensure the epoch falls within the plotting range
  1222. if clean_epoch_start >= start_idx and clean_epoch_end <= end_idx:
  1223. # Determine the run type for the epoch
  1224. clean_epoch_type = find_run_type(clean_epoch, folder_path_mouse_to_analyse, session_index, start_session_index=start_session_index)
  1225. # Assign color based on the run type
  1226. color = colors[run_types.index(clean_epoch_type)]
  1227. # Highlight cleaned run epochs with a semi-transparent overlay
  1228. ax.axvspan(
  1229. traject_time[clean_epoch_start], traject_time[clean_epoch_end],
  1230. color=color, alpha=0.3, linewidth=0,
  1231. label='Adjusted Run Epoch' if idx == 0 else "", zorder=2
  1232. )
  1233. # Highlight original run epochs with a slightly different transparency
  1234. ax.axvspan(
  1235. traject_time[epoch_start], traject_time[epoch_end],
  1236. color=color, alpha=0.3, linewidth=0,
  1237. label='Original Run Epoch' if idx == 0 else "", zorder=1
  1238. )
  1239. # Set plot limits and labels
  1240. ax.set_ylim(bottom=-5, top=max(speeds[start_idx:end_idx]) * 1.1)
  1241. ax.set_xlabel('Time (s)', fontsize=7)
  1242. ax.set_ylabel('Speed (cm/s)', fontsize=7)
  1243. # %% [markdown]
  1244. # ### Selecting the index of the session of which to show a portion of trajectory
  1245. # %%
  1246. index_example_session = 13
  1247. # %% [markdown]
  1248. # ### Selecting the portion of trajectory to show by epoch index (different mice can be chosen for panel A and panels B, C, D, E)
  1249. # %%
  1250. # Setting the path to the example mouse's data folder
  1251. folder_path_example_mouse_to_analyse = os.path.join(path_to_data_folder,example_mouse)
  1252. folder_path_example_mouse_bis_to_analyse = os.path.join(path_to_data_folder,example_mouse_bis)
  1253. # Find every sessions of the example mouse
  1254. example_sessions_to_analyse = sorted([name for name in os.listdir(folder_path_example_mouse_bis_to_analyse)
  1255. if os.path.isdir(os.path.join(folder_path_example_mouse_bis_to_analyse, name))
  1256. and name.startswith('MOU')])
  1257. # Select the name of the last session
  1258. example_session_to_analyse = example_sessions_to_analyse[first_and_last_session_indexes[0] + index_example_session]
  1259. # Load the session data from a pickle file
  1260. output_example_pickle_filename = f"{example_session_to_analyse}_basic_processing_output.pickle"
  1261. output_example_pickle_filepath = os.path.join(folder_path_example_mouse_bis_to_analyse, example_session_to_analyse, output_example_pickle_filename)
  1262. # Open and load the session data from the pickle file
  1263. with open(output_example_pickle_filepath, 'rb') as file:
  1264. session_data = pickle.load(file)
  1265. # Extract the time and speed of the mouse trajectory for this session
  1266. traject_time = session_data['timeofframes']
  1267. speeds = session_data['speeds']
  1268. # Define epochs
  1269. run_epochs, clean_run_epochs= detect_raw_and_clean_run_epochs(speeds,traject_time)
  1270. # %% [markdown]
  1271. # # Plot panel A
  1272. # This figure shows several epochs of the example session
  1273. # %%
  1274. # Select the index of the first and last epochs to show
  1275. first_epoch_to_plot = 10
  1276. last_epoch_to_plot = 16
  1277. # %%
  1278. if plotintermediatesteps:
  1279. fig=plt.figure(figsize=(cm2inch(18), cm2inch(5)), dpi=300, constrained_layout=False, facecolor='w')
  1280. gs = fig.add_gridspec(1, 1 , hspace=0.5)
  1281. ### Panel A ###
  1282. row1 = gs[0].subgridspec(1, 4, wspace=.3, hspace=.3, width_ratios=[1,1,1,1])
  1283. ax_11 = plt.subplot(row1[0],aspect="equal")
  1284. ax_12 = plt.subplot(row1[1:3])
  1285. ax_1bis = plt.subplot(row1[3])
  1286. start_index, end_index, trapeze_switch_times = plot_selected_run_epochs(folder_path_example_mouse_to_analyse, index_example_session, first_epoch_to_plot, last_epoch_to_plot, arena_coordinates_cm, ax_11, show_legend=True)
  1287. plot_trajectory_speed_chunk(start_index, end_index, traject_time, speeds, run_epochs, clean_run_epochs, ax_12, folder_path_example_mouse_to_analyse, index_example_session)
  1288. plot_learning_curves(mice_rewarded_qts_persession, ax_1bis, mice_to_highlight=illustration_mice, highlight_colors=illustration_colors, ylim=[0,300], show_individual_mice=True, median_color= 'violet', show_xlabel = False, xlim=xlim_stats, ylabel='QTs Nber', main_line_label='Rewarded QTs', tick_interval=2, index_offset=session_index_offset)
  1289. plot_learning_curves(mice_unrewarded_qts_persession, ax_1bis, mice_to_highlight=illustration_mice, highlight_colors=illustration_colors, ylim=[0,300], show_individual_mice=False, median_color= 'indigo', show_xlabel = False, xlim=xlim_stats, ylabel='QTs Nber', main_line_label='Unrewarded QTs', tick_interval=2, index_offset=session_index_offset)
  1290. # %% [markdown]
  1291. # # 9. Panel B, C
  1292. # %%
  1293. def plot_run_type(folder_path_mouse_to_analyse, session_index, arena_coordinates, ax, start_session_index=0, runtype='', q=4, time_start=None, time_end=None, show_legend=False, legend_loc = (1,1)):
  1294. """
  1295. Plots trajectories of the selected run types.
  1296. Arguments:
  1297. folder_path_mouse_to_analyse (str): Path to the folder containing session data.
  1298. session_index (int): Name of the session to analyse.
  1299. arena_coordinates (list): List of four 2D coordinates, locating the corners of the arena.
  1300. ax (matplotlib.axes.Axes): Matplotlib axis object for plotting.
  1301. start_session_index (int, optional): index of the first session in the list of sessions to analyse.
  1302. If different from 0, the session index should refer to
  1303. the index of the session in the sub-list of session to process, not the total list.
  1304. runtype (str): name of the type of run to plot. Choose among: "run_around_tower", "run_between_towers", "run_toward_tower", "exploratory_run", "immobility", "unclassified"
  1305. q (int, optional): Number of points used to determine movement direction to plot arrows. Default is 4.
  1306. time_start (float, optional): Time of the at which the trajectory will start being displayed
  1307. time_end (float, optional): Time of the at which the trajectory will stop being displayed
  1308. """
  1309. # Get all session folders that start with 'MOU' and sort them
  1310. sessions_to_analyse = sorted([name for name in os.listdir(folder_path_mouse_to_analyse)
  1311. if os.path.isdir(os.path.join(folder_path_mouse_to_analyse, name))
  1312. and name.startswith('MOU')])
  1313. session_to_analyse = sessions_to_analyse[start_session_index+session_index]
  1314. # Define the output pickle filename and its full path
  1315. output_pickle_filename = f"{session_to_analyse}_basic_processing_output.pickle"
  1316. output_pickle_filepath = os.path.join(folder_path_mouse_to_analyse, session_to_analyse, output_pickle_filename)
  1317. # Open and load the session data from the pickle file
  1318. with open(output_pickle_filepath, 'rb') as file:
  1319. session_data = pickle.load(file)
  1320. # Extract the time data for the trajectory
  1321. traject_time = session_data['timeofframes']
  1322. # Set default start and end times if not provided
  1323. if time_start is None:
  1324. time_start = traject_time[0]
  1325. if time_end is None:
  1326. time_end = traject_time[-1]
  1327. # Extract the smoothed X and Y positions of the mouse
  1328. smoothed_positions = session_data['positions']
  1329. smoothed_xpositions = smoothed_positions[0]
  1330. smoothed_ypositions = smoothed_positions[1]
  1331. # Extract epoch data and trapeze coordinates
  1332. all_epochs = session_data['all_epochs']
  1333. all_trapezes_coordinates_cm = session_data['all_trapezes_coordinates_cm']
  1334. # If runtype is not provided, raise a warning and return
  1335. if not runtype:
  1336. warnings.warn("The 'runtype' parameter is required and was not provided.")
  1337. return
  1338. # Retrieve the epochs corresponding to the provided runtype
  1339. runtype_epochs = filter_qts(all_epochs.get(runtype))
  1340. # If no epochs are found for the provided runtype, raise a warning and return
  1341. if runtype_epochs is None:
  1342. warnings.warn(f"The 'runtype' '{runtype}' is not found in 'all_epochs'.")
  1343. return
  1344. tower_coordinates = session_data['towers_coordinates_cm']
  1345. # Plot each tower
  1346. for i, (tower_name, vertices) in enumerate(tower_coordinates.items()):
  1347. tower_x, tower_y = zip(*vertices + [vertices[0]])
  1348. ax.fill(tower_x, tower_y, 'black', alpha=0.01)
  1349. ax.plot(tower_x, tower_y, 'k-', linewidth=0.5)
  1350. # Plot each trapeze
  1351. fill_colors = ['lightgray'] * 4 # Use a list of light blue colors for trapezes
  1352. for i, (tower, trapezes) in enumerate(all_trapezes_coordinates_cm.items()):
  1353. for j, (trapeze, coordinates) in enumerate(trapezes.items()):
  1354. # Close the trapeze polygon by appending the first vertex
  1355. coordinates_copy = coordinates + [coordinates[0]]
  1356. x_coords, y_coords = zip(*coordinates_copy)
  1357. # Fill the trapeze area with the color
  1358. ax.fill(x_coords, y_coords, color=fill_colors[j % len(fill_colors)], alpha=0.5)
  1359. displayed_labels = set([])
  1360. # Loop through each epoch in the runtype and plot the trajectory
  1361. for runtype_epoch in runtype_epochs:
  1362. # Get the start and end indices of the current epoch
  1363. start_index, end_index = runtype_epoch[0][0], runtype_epoch[0][1]
  1364. # Skip the epoch if it is outside the time window
  1365. if traject_time[start_index] < time_start or traject_time[end_index] > time_end:
  1366. continue
  1367. # Check that the start and end indices are within the bounds of the trajectory time array
  1368. if start_index < 0 or end_index >= len(traject_time):
  1369. print(f"Indexes out of bounds for runtype_epoch: {runtype_epoch}")
  1370. continue
  1371. # Extract the positions for the current epoch
  1372. runtype_epoch_xpositions = smoothed_xpositions[start_index:end_index + 1]
  1373. runtype_epoch_ypositions = smoothed_ypositions[start_index:end_index + 1]
  1374. numberofpositions = len(runtype_epoch_xpositions)
  1375. # Ensure that the lengths of X and Y positions match
  1376. if len(runtype_epoch_xpositions) != len(runtype_epoch_ypositions):
  1377. raise ValueError("The lengths of X and Y positions lists must be the same.")
  1378. # if the runtype is 'run_around_tower', check that it is not too long (above 2 seconds, i.e., 50 positions)
  1379. if runtype == 'run_around_tower' and numberofpositions > 50:
  1380. continue
  1381. # Loop through the positions to plot the trajectory
  1382. for index in range(numberofpositions - 2):
  1383. # Use a different color for run_around_tower if the direction matches the rewarded direction
  1384. if runtype == 'run_around_tower':
  1385. color = 'blue' if runtype_epoch[3]['direction'] == 'CW' else 'red'
  1386. # Label the lines for first occurrence of each color (CW,CCW)
  1387. if color == 'blue' and not('blue' in displayed_labels):
  1388. line_label = 'CW QTs'
  1389. displayed_labels.add('blue')
  1390. elif color == 'red' and not('red' in displayed_labels):
  1391. line_label = 'CCW QTs'
  1392. displayed_labels.add('red')
  1393. else:
  1394. line_label = None
  1395. else:
  1396. color = 'black'
  1397. # Plot the line segment between consecutive positions
  1398. ax.plot(runtype_epoch_xpositions[index:index + 2], runtype_epoch_ypositions[index:index + 2], color=color, linewidth=0.5, label=line_label)
  1399. # Plot the start point as a small black circle
  1400. ax.plot(runtype_epoch_xpositions[0], runtype_epoch_ypositions[0], color='black', marker='o', markersize=1)
  1401. # Compute the direction of the arrow using the last 'q' positions (default is 4)
  1402. if len(runtype_epoch_xpositions) >= q:
  1403. dx = runtype_epoch_xpositions[-1] - runtype_epoch_xpositions[-q]
  1404. dy = runtype_epoch_ypositions[-1] - runtype_epoch_ypositions[-q]
  1405. # Normalize the direction vector
  1406. norm = np.hypot(dx, dy)
  1407. if norm != 0:
  1408. dx /= norm
  1409. dy /= norm
  1410. # Plot the direction arrow at the end point
  1411. ax.arrow(runtype_epoch_xpositions[-1], runtype_epoch_ypositions[-1], dx, dy,
  1412. head_width=1, head_length=1, fc='black', ec='black')
  1413. # Draw the arena perimeter
  1414. arena_x, arena_y = zip(*arena_coordinates + [arena_coordinates[0]])
  1415. ax.plot(arena_x, arena_y, 'grey', linewidth=1)
  1416. # Hide the spines (axes) from the plot
  1417. for spine in ax.spines.values():
  1418. spine.set_visible(False)
  1419. # Remove the ticks from the x and y axes
  1420. ax.set_xticks([])
  1421. ax.set_yticks([])
  1422. if show_legend:
  1423. ax.legend(loc=legend_loc, frameon=False, ncols=2)
  1424. # %%
  1425. if plotintermediatesteps:
  1426. fig=plt.figure(figsize=(cm2inch(18), cm2inch(5)), dpi=300, constrained_layout=False, facecolor='w')
  1427. gs = fig.add_gridspec(1, 1 , hspace=0.5)
  1428. ### Panel B and C ###
  1429. row1 = gs[0].subgridspec(1, len(illustration_sessions_indexes)+1, wspace=.3, hspace=.1)
  1430. for j in range(len(illustration_sessions_indexes)):
  1431. ax_1 = plt.subplot(row1[j], aspect="equal")
  1432. plot_run_type(folder_path_example_mouse_to_analyse, illustration_sessions_indexes[j], arena_coordinates_cm, ax_1, runtype='run_around_tower', q=4)
  1433. day, period = get_day_and_period(illustration_sessions_indexes[j])
  1434. ax_1.text(0.5, 1.15, f'Session {illustration_sessions_indexes[j] + 1}', va='center', ha='center', transform=ax_1.transAxes, fontsize=7)
  1435. ax_1.text(0.5, 1.1 - 0.05, f'Day {day} ({period})', va='center', ha='center', transform=ax_1.transAxes, fontsize=5)
  1436. if j==0:
  1437. ax_1.text(-0.1, 0.5, f'Mouse {example_mouse_index+1}', color='darkorange',rotation=90, va='center', ha='center', transform=ax_1.transAxes, fontsize=7)
  1438. ax_1bis = plt.subplot(row1[len(illustration_sessions_indexes)])
  1439. plot_learning_curves(mice_ccw_vs_cw_norm_diff_persession, ax_1bis, mice_to_highlight=illustration_mice, highlight_colors=illustration_colors, show_xlabel = False, xlim=xlim_stats, ylim=[-1,1], ylabel='CCW-CW / CW+CCW', tick_interval=2, index_offset=session_index_offset, show_legend=False)
  1440. ax_1bis.hlines(0,0,15,color='grey',linestyle='--')
  1441. # %% [markdown]
  1442. # # 10. Panel E, F
  1443. # %%
  1444. def plot_runs_around_towers_origin(folder_path_mouse_to_analyse, session_index, ax, start_session_index=0, q=4, time_start=None, time_end=None, show_legend=True, xlim=None, ylim=None, show_xlabel=True, show_ylabel=True):
  1445. """
  1446. Plots the trajectory of the "run around tower" epochs, with the origin of each run aligned to its start position.
  1447. The plot also displays the direction of movement and optionally includes a legend, axis labels, and axis limits.
  1448. Arguments:
  1449. folder_path_mouse_to_analyse (str): Path to the folder containing the session's data.
  1450. session_index (int): Index of the session to analyse.
  1451. ax (matplotlib.axes.Axes): The axis object where the plot will be drawn.
  1452. start_session_index (int, optional): index of the first session in the list of sessions to analyse.
  1453. If different from 0, the session index should refer to
  1454. the index of the session in the sub-list of session to process, not the total list.
  1455. q (int, optional): Number of previous points used to compute the direction arrow. Defaults to 4.
  1456. time_start (float, optional): The start time for plotting. Defaults to the first time in the data.
  1457. time_end (float, optional): The end time for plotting. Defaults to the last time in the data.
  1458. show_legend (bool, optional): Whether to show the legend indicating rewarded and unrewarded directions. Defaults to True.
  1459. xlim (tuple, optional): Limits for the x-axis.
  1460. ylim (tuple, optional): Limits for the y-axis.
  1461. show_xlabel (bool, optional): Whether to show the x-axis label. Defaults to True.
  1462. show_ylabel (bool, optional): Whether to show the y-axis label. Defaults to True.
  1463. """
  1464. # Get all session folders that start with 'MOU' and sort them
  1465. sessions_to_analyse = sorted([name for name in os.listdir(folder_path_mouse_to_analyse)
  1466. if os.path.isdir(os.path.join(folder_path_mouse_to_analyse, name))
  1467. and name.startswith('MOU')])
  1468. session_to_analyse = sessions_to_analyse[start_session_index+session_index]
  1469. # Define the output pickle filename and its full path
  1470. output_pickle_filename = f"{session_to_analyse}_basic_processing_output.pickle"
  1471. output_pickle_filepath = os.path.join(folder_path_mouse_to_analyse, session_to_analyse, output_pickle_filename)
  1472. # Open and load the session data from the pickle file
  1473. with open(output_pickle_filepath, 'rb') as file:
  1474. session_data = pickle.load(file)
  1475. # Extract the trajectory time information
  1476. traject_time = session_data['timeofframes']
  1477. # Set default start and end times if not provided
  1478. if time_start is None:
  1479. time_start = traject_time[0]
  1480. if time_end is None:
  1481. time_end = traject_time[-1]
  1482. # Extract the smoothed positions (X and Y) from the session data
  1483. smoothed_positions = session_data['positions']
  1484. smoothed_xpositions = smoothed_positions[0]
  1485. smoothed_ypositions = smoothed_positions[1]
  1486. # Retrieve the epochs corresponding to 'run_around_tower'
  1487. runs_around_tower = filter_qts(copy.deepcopy(session_data['all_epochs']['run_around_tower']))
  1488. # Define the fixed origin at (0, 0)
  1489. fixed_origin = (0, 0)
  1490. # Loop through each 'run_around_tower' epoch to plot the trajectory
  1491. for run_around_tower in runs_around_tower:
  1492. start_index, end_index = run_around_tower[0][0], run_around_tower[0][1]
  1493. # Skip epochs that fall outside the specified time range
  1494. if traject_time[start_index] < time_start or traject_time[end_index] > time_end:
  1495. continue
  1496. # Extract the X and Y positions for the current run
  1497. runtype_epoch_xpositions = smoothed_xpositions[start_index:end_index + 1]
  1498. runtype_epoch_ypositions = smoothed_ypositions[start_index:end_index + 1]
  1499. numberofpositions = len(runtype_epoch_xpositions)
  1500. # check that the detected run is not too long (above 2 seconds, i.e., 50 positions)
  1501. if numberofpositions > 50:
  1502. continue
  1503. # Determine the start position and translate the coordinates to have this as the origin
  1504. start_x, start_y = runtype_epoch_xpositions[0], runtype_epoch_ypositions[0]
  1505. translated_xpositions = [x - start_x + fixed_origin[0] for x in runtype_epoch_xpositions]
  1506. translated_ypositions = [y - start_y + fixed_origin[1] for y in runtype_epoch_ypositions]
  1507. # Loop through positions and plot the trajectory between consecutive points
  1508. for i in range(numberofpositions - 1):
  1509. # Choose color based on direction
  1510. color = 'blue' if run_around_tower[3]['direction'] == 'CW' else 'red'
  1511. # Plot the trajectory between two consecutive points
  1512. ax.plot(translated_xpositions[i:i+2], translated_ypositions[i:i+2], color=color, linewidth=0.5)
  1513. # Plot the start point as a black marker
  1514. ax.plot(translated_xpositions[0], translated_ypositions[0], marker='o', color='black', linewidth=0.5, markersize=1)
  1515. # Compute and plot an arrow showing the direction of movement based on the last 'q' points
  1516. if len(translated_xpositions) >= q:
  1517. dx = translated_xpositions[-1] - translated_xpositions[-q]
  1518. dy = translated_ypositions[-1] - translated_ypositions[-q]
  1519. norm_speed = np.hypot(dx, dy)
  1520. if norm_speed != 0:
  1521. dx /= norm_speed
  1522. dy /= norm_speed
  1523. # Plot the arrow in red to indicate movement direction
  1524. ax.arrow(translated_xpositions[-1], translated_ypositions[-1], dx, dy, head_width=1, head_length=1, fc='black', ec='black')
  1525. # Set the x-axis label if specified
  1526. if show_xlabel:
  1527. ax.set_xlabel('X Position (cm)')
  1528. # Set the y-axis label if specified
  1529. if show_ylabel:
  1530. ax.set_ylabel('Y Position (cm)')
  1531. # Set the limits for the x and y axes if provided
  1532. ax.set_xlim(xlim)
  1533. ax.set_ylim(ylim)
  1534. # Hide the top and right spines for better aesthetics
  1535. ax.spines['right'].set_visible(False)
  1536. ax.spines['top'].set_visible(False)
  1537. # Show the legend if specified
  1538. if show_legend:
  1539. ax.legend(loc=[0.4, 1.2], frameon=False)
  1540. # %%
  1541. if plotintermediatesteps:
  1542. fig=plt.figure(figsize=(cm2inch(18), cm2inch(5)), dpi=300, constrained_layout=False, facecolor='w')
  1543. gs = fig.add_gridspec(1, 1 , hspace=0.5)
  1544. ### Panel C ###
  1545. row2 = gs[0].subgridspec(1, len(illustration_sessions_indexes)+1, wspace=.3, hspace=.3)
  1546. for k in range(len(illustration_sessions_indexes)):
  1547. ax_2 = plt.subplot(row2[k], aspect="equal")
  1548. plot_runs_around_towers_origin(folder_path_example_mouse_to_analyse, illustration_sessions_indexes[k], ax_2, q=4, show_legend= True if k==0 else False, show_xlabel= True if k==0 else False, show_ylabel= True if k==0 else False, xlim=(-30,30), ylim=(-30,30))
  1549. ax_2.set_xticks([-20, 0, 20])
  1550. ax_2.set_yticks([-20, 0, 20])
  1551. ax_2bis = plt.subplot(row2[len(illustration_sessions_indexes)])
  1552. plot_learning_curves(overall_cw_turns_frechet_distances_per_session, ax_2bis, mice_to_highlight=illustration_mice, highlight_colors=illustration_colors, show_individual_mice=False, xlim=xlim_stats, median_color='blue', show_xlabel = True, ylabel='Fréchet distance (cm)', tick_interval=2, index_offset=session_index_offset, show_legend=False)
  1553. plot_learning_curves(overall_ccw_turns_frechet_distances_per_session, ax_2bis, mice_to_highlight=illustration_mice, highlight_colors=illustration_colors, show_individual_mice=True, xlim=xlim_stats, median_color='red', show_xlabel = True, ylabel='Fréchet distance (cm)', tick_interval=2, index_offset=session_index_offset, show_legend=False)
  1554. # %% [markdown]
  1555. # # 11. Panel G, H
  1556. # %%
  1557. def plot_runs_around_speed_profiles(folder_path_mouse_to_analyse, session_index, ax, start_session_index=0, time_start=None, time_end=None, xlim=None, ylim=None, show_xlabel=True, show_ylabel=True, show_color_bar=True):
  1558. """
  1559. This function plots the speed profiles of "runs around tower" epochs from a session's data.
  1560. Arguments:
  1561. folder_path_mouse_to_analyse (str): The directory path containing the session's data.
  1562. session_index (str): Index of the session to analyse.
  1563. ax (matplotlib.axes.Axes): The axis object where the plot will be drawn.
  1564. start_session_index (int, optional): index of the first session in the list of sessions to analyse.
  1565. If different from 0, the session index should refer to
  1566. the index of the session in the sub-list of session to process, not the total list.
  1567. time_start (float, optional): The start time for the plot. Defaults to the first time point.
  1568. time_end (float, optional): The end time for the plot. Defaults to the last time point.
  1569. xlim (tuple, optional): Limits for the x-axis.
  1570. ylim (tuple, optional): Limits for the y-axis.
  1571. show_xlabel (bool, optional): Whether to show the x-axis label. Defaults to True.
  1572. show_ylabel (bool, optional): Whether to show the y-axis label. Defaults to True.
  1573. show_color_bar (bool, optional): Whether to display a color bar. Defaults to True.
  1574. """
  1575. # Get all session folders that start with 'MOU' and sort them
  1576. sessions_to_analyse = sorted([name for name in os.listdir(folder_path_mouse_to_analyse)
  1577. if os.path.isdir(os.path.join(folder_path_mouse_to_analyse, name))
  1578. and name.startswith('MOU')])
  1579. session_to_analyse = sessions_to_analyse[start_session_index+session_index]
  1580. # Define the output pickle filename and its full path
  1581. output_pickle_filename = f"{session_to_analyse}_basic_processing_output.pickle"
  1582. output_pickle_filepath = os.path.join(folder_path_mouse_to_analyse, session_to_analyse, output_pickle_filename)
  1583. # Open and load the session data from the pickle file
  1584. with open(output_pickle_filepath, 'rb') as file:
  1585. session_data = pickle.load(file)
  1586. # Extract the trajectory time information
  1587. traject_time = session_data['timeofframes']
  1588. # Set default start and end times if not provided
  1589. if time_start is None:
  1590. time_start = traject_time[0]
  1591. if time_end is None:
  1592. time_end = traject_time[-1]
  1593. # Extract the smoothed positions (X and Y) and speed data from session
  1594. smoothed_positions = session_data['positions']
  1595. smoothed_xpositions = smoothed_positions[0]
  1596. smoothed_ypositions = smoothed_positions[1]
  1597. speeds = session_data['speeds']
  1598. # Retrieve the epochs corresponding to 'run_around_tower' #filter QT keep only real quarter of turns
  1599. runs_around_tower = filter_qts(copy.deepcopy(session_data['all_epochs']['run_around_tower']))
  1600. # Initialize a counter for the total number of runs within the specified time window
  1601. n_total = 0
  1602. # Loop through each 'run_around_tower' epoch and count how many are within the time window
  1603. for run_around_tower in runs_around_tower:
  1604. if run_around_tower[3]['direction'] == 'CW':continue # Skip if the direction is not 'CCW' (i.e., only consider counter-clockwise runs)
  1605. start_index, end_index = run_around_tower[0][0], run_around_tower[0][1]
  1606. # Skip the epochs that fall outside the specified time range
  1607. if traject_time[start_index] < time_start or traject_time[end_index] > time_end:
  1608. continue
  1609. # Increment the total count of runs within the time window
  1610. n_total += 1
  1611. # Define colormap for color representation of the runs
  1612. cmap = plt.cm.viridis
  1613. norm = Normalize(vmin=0, vmax=n_total) # Normalize for color mapping
  1614. # Local index to differentiate each run in the color scale
  1615. local_index = 0
  1616. # Loop through the 'run_around_tower' epochs again to plot the speed profile
  1617. for run_around_tower in runs_around_tower:
  1618. # if run_around_tower[3]['direction'] == 'CW':continue # Skip if the direction is not 'CCW' (i.e., only consider counter-clockwise runs)
  1619. start_index, end_index = run_around_tower[0][0], run_around_tower[0][1]
  1620. # Skip the epochs outside the time window
  1621. if traject_time[start_index] < time_start or traject_time[end_index] > time_end:
  1622. continue
  1623. # Extract the positions for the current run
  1624. runtype_epoch_xpositions = smoothed_xpositions[start_index:end_index + 1]
  1625. runtype_epoch_ypositions = smoothed_ypositions[start_index:end_index + 1]
  1626. # Ensure that X and Y positions lists have the same length
  1627. if len(runtype_epoch_xpositions) != len(runtype_epoch_ypositions):
  1628. raise ValueError("The lengths of X and Y positions lists must be the same.")
  1629. # Check if the run is too long (above 2 seconds, i.e., 50 positions)
  1630. if len(runtype_epoch_xpositions) > 50:
  1631. continue
  1632. # Adjust the time to start from zero for the current run epoch
  1633. adjusted_time = [t - traject_time[start_index] for t in traject_time[start_index:end_index + 1]]
  1634. # Plot the speed for the current run
  1635. ax.plot(adjusted_time, speeds[start_index:end_index + 1], color=cmap(norm(local_index)), alpha=0.5)
  1636. # Increment the local index for color mapping
  1637. local_index += 1
  1638. # Set the x-axis label if specified
  1639. if show_xlabel:
  1640. ax.set_xlabel('Time (s)')
  1641. # Set the y-axis label if specified
  1642. if show_ylabel:
  1643. ax.set_ylabel('Speed (cm/s)')
  1644. # Set the limits for the x and y axes if provided
  1645. ax.set_xlim(xlim)
  1646. ax.set_ylim(ylim)
  1647. # Add color bar if specified
  1648. if show_color_bar:
  1649. cbax = ax.inset_axes([0.25, 0.95, 0.5, 0.07]) # Define the position for the color bar
  1650. cbar = fig.colorbar(plt.cm.ScalarMappable(cmap=cmap, norm=norm), cax=cbax, orientation='horizontal')
  1651. cbar.set_label('Index of Run', fontsize=5)
  1652. cbar.set_ticks([0, n_total]) # Set color bar ticks
  1653. cbar.set_ticklabels(['First', 'Last'], fontsize=5) # Label the ticks as First and Last
  1654. # %%
  1655. if plotintermediatesteps:
  1656. fig=plt.figure(figsize=(cm2inch(18), cm2inch(5)), dpi=300, constrained_layout=False, facecolor='w')
  1657. gs = fig.add_gridspec(1, 1 , hspace=0.5)
  1658. ### Panel D ###
  1659. row3 = gs[0].subgridspec(1, len(illustration_sessions_indexes)+1, wspace=.3, hspace=.3)
  1660. for l in range(len(illustration_sessions_indexes)):
  1661. ax_3 = plt.subplot(row3[l], aspect="equal")
  1662. plot_runs_around_speed_profiles(folder_path_example_mouse_to_analyse, illustration_sessions_indexes[l], ax_3, xlim=(0,2), ylim=(0,80), show_xlabel= True if l==0 else False, show_ylabel= True if l==0 else False, show_color_bar= True if l==0 else False)
  1663. force_aspect(ax_3,ratio=1)
  1664. ax_3bis = plt.subplot(row3[len(illustration_sessions_indexes)])
  1665. plot_learning_curves(mice_median_maximum_cw_turn_speed_persession, ax_3bis, mice_to_highlight=illustration_mice, highlight_colors=illustration_colors, show_individual_mice=False, median_color= 'blue', show_xlabel = False, xlim=xlim_stats, ylabel='Peak speed (cm/s)', main_line_label="CW QTs", tick_interval=2, index_offset=session_index_offset, show_legend=False)
  1666. plot_learning_curves(mice_median_maximum_ccw_turn_speed_persession, ax_3bis, mice_to_highlight=illustration_mice, highlight_colors=illustration_colors, show_individual_mice=True, median_color= 'red', show_xlabel = False, xlim=xlim_stats, ylabel='Peak speed (cm/s)', main_line_label="CCW QTs", tick_interval=2, index_offset=session_index_offset, show_legend=True)
  1667. # %% [markdown]
  1668. # # 12. Panel I, J
  1669. # %%
  1670. def plot_cumulated_turns_time_profile(folder_path_mouse_to_analyse, session_to_analyse, ax, time_start=None, time_end=None, xlim=None, ylim=None, show_xlabel=True,
  1671. show_ylabel=True, show_legend=True, legend_loc=[1,1]):
  1672. """
  1673. This function plots the cumulative number of turns (CW vs. CCW) over time during a session.
  1674. Arguments:
  1675. folder_path_mouse_to_analyse (str): The directory where the mouse data is stored.
  1676. session_to_analyse (str): The name of the session to analyse.
  1677. ax (matplotlib.axes.Axes): The axes to plot the graph on.
  1678. time_start (float, optional): The start time for the plot. Defaults to the first time frame.
  1679. time_end (float, optional): The end time for the plot. Defaults to the last time frame.
  1680. xlim (tuple, optional): The limits for the x-axis (time). Defaults to None, which auto-scales.
  1681. ylim (tuple, optional): The limits for the y-axis (cumulative number of turns). Defaults to None, which auto-scales.
  1682. show_xlabel (bool, optional): Whether to display the x-axis label. Defaults to True.
  1683. show_ylabel (bool, optional): Whether to display the y-axis label. Defaults to True.
  1684. show_legend (bool, optional): Whether to display the legend. Defaults to True.
  1685. """
  1686. # Load the session data from a pickle file
  1687. output_pickle_filename = f"{session_to_analyse}_basic_processing_output.pickle"
  1688. output_pickle_filepath = os.path.join(folder_path_mouse_to_analyse, session_to_analyse, output_pickle_filename)
  1689. with open(output_pickle_filepath, 'rb') as file:
  1690. session_data = pickle.load(file)
  1691. # Extract trajectory time and check if the time range is provided
  1692. traject_time = session_data['timeofframes']
  1693. if time_start is None:
  1694. time_start = traject_time[0]
  1695. if time_end is None:
  1696. time_end = traject_time[-1]
  1697. # Extract the runs around tower data
  1698. runs_around_tower = filter_qts(copy.deepcopy(session_data['all_epochs']['run_around_tower']))
  1699. # Prepare lists for CW and CCW times
  1700. time_of_runsaroundtower_cw = []
  1701. time_of_runsaroundtower_ccw = []
  1702. for run_around_tower in runs_around_tower:
  1703. start_index, end_index = run_around_tower[0][0], run_around_tower[0][1]
  1704. # Check if the current run is within the time window
  1705. if traject_time[start_index] < time_start or traject_time[end_index] > time_end:
  1706. continue
  1707. # Separate times based on direction (CW or CCW)
  1708. if run_around_tower[3]['direction'] == 'CW':
  1709. time_of_runsaroundtower_cw.append(run_around_tower[4]['epoch_time'])
  1710. if run_around_tower[3]['direction'] == 'CCW':
  1711. time_of_runsaroundtower_ccw.append(run_around_tower[4]['epoch_time'])
  1712. # Sort CW and CCW times
  1713. cw_times_sorted = np.sort(time_of_runsaroundtower_cw)
  1714. ccw_times_sorted = np.sort(time_of_runsaroundtower_ccw)
  1715. # Calculate the cumulative counts
  1716. cw_cumulative = np.arange(1, len(cw_times_sorted) + 1)
  1717. ccw_cumulative = np.arange(1, len(ccw_times_sorted) + 1)
  1718. # Plot the cumulative number of good and bad turns over time
  1719. ax.step(cw_times_sorted, cw_cumulative, where='post', label='CW QTs', color='blue')
  1720. ax.step(ccw_times_sorted, ccw_cumulative, where='post', label='CCW QTs', color='red')
  1721. # Set axis labels and limits based on user input
  1722. if show_xlabel:
  1723. ax.set_xlabel('Session Time (s)')
  1724. if show_ylabel:
  1725. ax.set_ylabel('Cum. nber of QTs')
  1726. ax.set_xlim(xlim)
  1727. ax.set_ylim(ylim)
  1728. # Display the legend if requested
  1729. if show_legend:
  1730. ax.legend(ncols=2, loc=legend_loc, frameon=False)
  1731. # Hide unnecessary spines for cleaner visualization
  1732. ax.spines['right'].set_visible(False)
  1733. ax.spines['top'].set_visible(False)
  1734. # %%
  1735. if plotintermediatesteps:
  1736. fig=plt.figure(figsize=(cm2inch(18), cm2inch(5)), dpi=300, constrained_layout=False, facecolor='w')
  1737. gs = fig.add_gridspec(1, 1 , hspace=0.5)
  1738. ### Panel E ###
  1739. row4 = gs[0].subgridspec(1, len(illustration_sessions_indexes)+1, wspace=.3, hspace=.3)
  1740. for m in range(len(illustration_sessions_indexes)):
  1741. ax_4 = plt.subplot(row4[m], aspect="equal")
  1742. show_xlabel= True if m==0 else False
  1743. show_ylabel = True if m==0 else False
  1744. xlim = [0,600]
  1745. ylim = [0,250]
  1746. plot_cumulated_turns_time_profile(folder_path_example_mouse_to_analyse, illustration_mice[illustration_sessions_indexes[m]], ax_4, xlim=xlim, ylim=ylim , show_xlabel=show_xlabel, show_ylabel=show_ylabel, show_legend= False)
  1747. force_aspect(ax_4,ratio=1)
  1748. ax_4bis = plt.subplot(row4[len(illustration_sessions_indexes)])
  1749. plot_learning_curves(mice_qts_rewarded_dir_threshold_persession, ax_4bis, mice_to_highlight=illustration_mice, highlight_colors=illustration_colors, show_xlabel = True, xlim=xlim_stats, ylabel='Time to 80% of rewarded QTs (s)', tick_interval=2, index_offset=session_index_offset, show_legend=False)
  1750. # %% [markdown]
  1751. # # 13. Whole figure
  1752. # %%
  1753. fig=plt.figure(figsize=(cm2inch(15), cm2inch(20)), dpi=300, constrained_layout=False, facecolor='w')
  1754. gs = fig.add_gridspec(5, 2 , hspace=0.3, wspace=0.3, width_ratios=(1,0.5))
  1755. axes_to_align = []
  1756. axes_bis_to_align = []
  1757. ### Panel A, C ###
  1758. row1 = gs[0,0].subgridspec(1, 3, wspace=.5, hspace=.3, width_ratios=[1,1,1])
  1759. ax_11 = plt.subplot(row1[0],aspect="equal")
  1760. ax_12 = plt.subplot(row1[1:])
  1761. row1bis = gs[0,1].subgridspec(1, 1)
  1762. ax_1bis = plt.subplot(row1bis[0])
  1763. start_index, end_index, trapeze_switch_times = plot_selected_run_epochs(folder_path_example_mouse_bis_to_analyse, index_example_session, first_epoch_to_plot, last_epoch_to_plot, arena_coordinates_cm, ax_11, show_legend=True)
  1764. plot_trajectory_speed_chunk(start_index, end_index, traject_time, speeds, run_epochs, clean_run_epochs, ax_12, folder_path_example_mouse_bis_to_analyse, index_example_session)
  1765. ax_12.set_xlabel('Time (s)', labelpad=0)
  1766. axes_to_align.append(ax_11)
  1767. plot_learning_curves(mice_ccw_vs_cw_norm_diff_persession, ax_1bis, mice_to_highlight=illustration_mice, highlight_colors=illustration_colors, mice_to_highlight_labels=['Mouse 1', 'Mouse 2'], show_xlabel = False, xlim=xlim_stats, ylim=[-1,1], ylabel='CCW-CW / CW+CCW', tick_interval=2, index_offset=session_index_offset, main_line_label=f'n={len(mice_to_analyse)}', show_legend=True, legend_loc=(0.6,0.1))
  1768. plot_shuffled_spearman_test_res(ax_1bis, (4,-0.8), mice_ccw_vs_cw_norm_diff_persession, 1000, 'increasing', color='black', first_and_last_session_indexes=[1,14])
  1769. ax_1bis.axhline(0,linestyle='--', color='grey')
  1770. axes_bis_to_align.append(ax_1bis)
  1771. fig.text(.06, 0.89, 'A', weight='bold', va='center', ha='center', fontsize=7)
  1772. fig.text(.59, 0.89, 'C', weight='bold', va='center', ha='center', fontsize=7)
  1773. ### Panel B, D ###
  1774. row2 = gs[1,0].subgridspec(1, len(illustration_sessions_indexes), wspace=.3, hspace=.1)
  1775. row2bis = gs[1,1].subgridspec(1, 1)
  1776. ax_2bis = plt.subplot(row2bis[0])
  1777. for j in range(len(illustration_sessions_indexes)):
  1778. ax_2 = plt.subplot(row2[j], aspect="equal")
  1779. plot_run_type(folder_path_example_mouse_to_analyse, illustration_sessions_indexes[j], arena_coordinates_cm, ax_2, runtype='run_around_tower', q=4, show_legend= False, legend_loc=(0,-0.2))
  1780. day, period = get_day_and_period(illustration_sessions_indexes[j])
  1781. ax_2.text(0.5, 1.15, f'Session {illustration_sessions_indexes[j] + 1}', va='center', ha='center', transform=ax_2.transAxes, fontsize=7)
  1782. ax_2.text(0.5, 1.1 - 0.05, f'Day {day} ({period})', va='center', ha='center', transform=ax_2.transAxes, fontsize=5)
  1783. if j==0:
  1784. axes_to_align.append(ax_2)
  1785. ax_2.text(-0.38, 0.5, f'Mouse {example_mouse_index+1}', color=illustration_colors[example_mouse_index],rotation=90, va='center', ha='center', transform=ax_2.transAxes, fontsize=7)
  1786. plot_learning_curves(mice_rewarded_qts_persession, ax_2bis, mice_to_highlight=illustration_mice, highlight_colors=illustration_colors, show_individual_mice=True, median_color= 'violet', show_xlabel = False, xlim=xlim_stats, ylim=[0,300], ylabel='Total nber of QTs', main_line_label=f'Rewarded', tick_interval=2, index_offset=session_index_offset, show_legend=True)
  1787. plot_learning_curves(mice_unrewarded_qts_persession, ax_2bis, show_individual_mice=False, median_color= 'indigo', show_xlabel = False, xlim=xlim_stats, ylim=[0,300], ylabel='Total nber of QTs', main_line_label=f'Unrewarded', tick_interval=2, index_offset=session_index_offset, show_legend=True, legend_loc=(0.2, 0.81))
  1788. plot_shuffled_spearman_test_res(ax_2bis, (2.4,200), mice_rewarded_qts_persession, 1000, 'increasing', color='violet', first_and_last_session_indexes=[2,14])
  1789. plot_shuffled_spearman_test_res(ax_2bis, (2.4,150), mice_unrewarded_qts_persession, 1000, 'increasing', color='indigo', first_and_last_session_indexes=[2,14])
  1790. axes_bis_to_align.append(ax_2bis)
  1791. fig.text(.06, 0.72, 'B', weight='bold', va='center', ha='center', fontsize=7)
  1792. fig.text(.59, 0.72, 'D', weight='bold', va='center', ha='center', fontsize=7)
  1793. ### Panel E, F ###
  1794. row3 = gs[2,0].subgridspec(1, len(illustration_sessions_indexes), wspace=.3, hspace=.3)
  1795. sessions_to_analyse = sorted([name for name in os.listdir(folder_path_example_mouse_to_analyse)
  1796. if os.path.isdir(os.path.join(folder_path_example_mouse_to_analyse, name))
  1797. and name.startswith('MOU')])
  1798. for m in range(len(illustration_sessions_indexes)):
  1799. ax_3 = plt.subplot(row3[m], aspect="equal")
  1800. show_xlabel= True if m==0 else False
  1801. show_ylabel = True if m==0 else False
  1802. xlim = [0,600]
  1803. ylim = [0,320]
  1804. plot_cumulated_turns_time_profile(folder_path_example_mouse_to_analyse, sessions_to_analyse[illustration_sessions_indexes[m]], ax_3, xlim=xlim, ylim=[0,200] , show_xlabel=show_xlabel, show_ylabel=show_ylabel, show_legend= True if m==0 else False, legend_loc=(0,1.4))
  1805. force_aspect(ax_3,ratio=1)
  1806. if m==0:
  1807. axes_to_align.append(ax_3)
  1808. row3bis = gs[2,1].subgridspec(1, 1)
  1809. ax_3bis = plt.subplot(row3bis[0])
  1810. plot_learning_curves(mice_qts_rewarded_dir_threshold_persession, ax_3bis, mice_to_highlight=illustration_mice, highlight_colors=illustration_colors, show_xlabel = False, xlim=xlim_stats,ylim=[100,800] ,ylabel='Time to 80%\n of CCW QTs (s)', tick_interval=2, index_offset=session_index_offset, show_legend=False)
  1811. plot_shuffled_spearman_test_res(ax_3bis, (9,550), mice_qts_rewarded_dir_threshold_persession, 1000, 'decreasing', color='black', first_and_last_session_indexes=[2,14])
  1812. axes_bis_to_align.append(ax_3bis)
  1813. fig.text(.06, 0.57, 'E', weight='bold', va='center', ha='center', fontsize=7)
  1814. fig.text(.59, 0.57, 'F', weight='bold', va='center', ha='center', fontsize=7)
  1815. ### Panel G, H ###
  1816. row4 = gs[3,0].subgridspec(1, len(illustration_sessions_indexes), wspace=.3, hspace=.3)
  1817. for k in range(len(illustration_sessions_indexes)):
  1818. ax_4 = plt.subplot(row4[k], aspect="equal")
  1819. plot_runs_around_towers_origin(folder_path_example_mouse_to_analyse, illustration_sessions_indexes[k], ax_4, q=4, show_xlabel= True if k==0 else False, show_ylabel= True if k==0 else False, show_legend= False, xlim=(-20,20), ylim=(-20,20))
  1820. ax_4.set_xticks([-20, 0, 20])
  1821. ax_4.set_yticks([-20, 0, 20])
  1822. if k==0:
  1823. axes_to_align.append(ax_4)
  1824. row4bis = gs[3,1].subgridspec(1, 1)
  1825. ax_4bis = plt.subplot(row4bis[0])
  1826. plot_learning_curves(overall_cw_turns_frechet_distances_per_session, ax_4bis, mice_to_highlight=illustration_mice, highlight_colors=illustration_colors, show_individual_mice=False, xlim=xlim_stats, median_color='blue', main_line_label='CW QTs', show_xlabel = False, ylabel='QTs Fréchet dist. (cm)', tick_interval=2, index_offset=session_index_offset, show_legend=False)
  1827. plot_learning_curves(overall_ccw_turns_frechet_distances_per_session, ax_4bis, mice_to_highlight=illustration_mice, highlight_colors=illustration_colors, show_individual_mice=True, xlim=xlim_stats, median_color='red', main_line_label='CCW QTs', show_xlabel = False, ylabel='QTs Fréchet dist. (cm)', tick_interval=2, index_offset=session_index_offset, show_legend=True, legend_loc=(0.4, 0.73))
  1828. plot_shuffled_spearman_test_res(ax_4bis, (10,3.5), overall_ccw_turns_frechet_distances_per_session, 1000, 'decreasing', color='red', first_and_last_session_indexes=[1,14])
  1829. plot_shuffled_spearman_test_res(ax_4bis, (6,8), overall_cw_turns_frechet_distances_per_session, 1000, 'decreasing', color='blue', first_and_last_session_indexes=[1,14])
  1830. axes_bis_to_align.append(ax_4bis)
  1831. fig.text(.06, 0.4, 'G', weight='bold', va='center', ha='center', fontsize=7)
  1832. fig.text(.59, 0.4, 'H', weight='bold', va='center', ha='center', fontsize=7)
  1833. ### Panel I, J ###
  1834. row5 = gs[4,0].subgridspec(1, len(illustration_sessions_indexes), wspace=.3, hspace=.3)
  1835. for l in range(len(illustration_sessions_indexes)):
  1836. ax_5 = plt.subplot(row5[l], aspect="equal")
  1837. plot_runs_around_speed_profiles(folder_path_example_mouse_to_analyse, illustration_sessions_indexes[l], ax_5, xlim=(0,1.5), ylim=(0,80), show_xlabel= True if l==0 else False, show_ylabel= True if l==0 else False, show_color_bar= True if l==0 else False)
  1838. if l==0:
  1839. axes_to_align.append(ax_5)
  1840. force_aspect(ax_5,ratio=1)
  1841. row5bis = gs[4,1].subgridspec(1, 1)
  1842. ax_5bis = plt.subplot(row5bis[0])
  1843. plot_learning_curves(mice_median_maximum_cw_turn_speed_persession, ax_5bis, mice_to_highlight=illustration_mice, highlight_colors=illustration_colors, show_individual_mice=False, median_color= 'blue', show_xlabel = False, xlim=xlim_stats, ylabel='QTs peak speed (cm/s)', main_line_label="CW QTs", tick_interval=2, index_offset=session_index_offset, show_legend=False)
  1844. plot_learning_curves(mice_median_maximum_ccw_turn_speed_persession, ax_5bis, mice_to_highlight=illustration_mice, highlight_colors=illustration_colors, show_individual_mice=True, median_color= 'red', show_xlabel = True, xlim=xlim_stats, ylabel='QTs peak speed (cm/s)', main_line_label="CCW QTs", tick_interval=2, index_offset=session_index_offset, show_legend=False)
  1845. plot_shuffled_spearman_test_res(ax_5bis, (2.6,45), mice_median_maximum_ccw_turn_speed_persession, 1000, 'increasing', color='red', first_and_last_session_indexes=[1,14])
  1846. plot_shuffled_spearman_test_res(ax_5bis, (7,15), mice_median_maximum_cw_turn_speed_persession, 1000, 'increasing', color='blue', first_and_last_session_indexes=[1,14])
  1847. axes_bis_to_align.append(ax_5bis)
  1848. fig.text(.06, 0.23, 'I', weight='bold', va='center', ha='center', fontsize=7)
  1849. fig.text(.59, 0.23, 'J', weight='bold', va='center', ha='center', fontsize=7)
  1850. fig.align_ylabels(axes_to_align)
  1851. fig.align_ylabels(axes_bis_to_align)
  1852. # After all plotting is done, right before saving:
  1853. fig.tight_layout()
  1854. plt.savefig("Figure03.png", facecolor='w', edgecolor='none', format="png", dpi=300)
  1855. # %%
  1856. # Save the figure as a PDF
  1857. fig.savefig("Figure03.pdf", format="pdf", bbox_inches='tight', dpi=300)

Figure03_group1_qt_example_and_stats.ipynb at commit e3bdaa3, no license · at the source

Overview

Authors: Maud Schaffhauser1, Tom Orjollet-Lacomme1, Kenza Amroune1, Thomas Morvan1, Aurélien Fortoul2, Mathias Lechelon3, David Robbe1
ORCID iDs: David Robbe
  1. INMED, INSERM, Aix-Marseille University, Turing Centre for Living Systems, Marseille, France
  2. INMED, INSERM, Aix-Marseille University, Marseille, France
  3. Aix Marseille University, CNRS, INSERM, MEP Centuri, Turing Center for Living Systems, Marseille, France
Journal: iScience, volume 29, issue 7, article 116498
Dates: received 28 October 2025; accepted 5 June 2026; published online 26 June 2026
Type: Research article · Language: English
License: CC BY-NC
Identifiers: DOI 10.1016/j.isci.2026.116498 · PMID 42491644 · PMCID PMC13377965 · OpenAlex W7166152200
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: mouse (organism)
Methods: Connectivity, Statistics, Graphs
Keywords: biological sciences, behavioral neuroscience, laboratory animal science
Topic: Zebrafish Biomedical Research Applications (Cell Biology, Biochemistry, Genetics and Molecular Biology), according to OpenAlex
Citations: not cited yet (Europe PMC); 99 references in the paper
Research resources: C57BL/6N wildtype mice RRID:IMSR_CRL:027, C57BL/6J DRD1-Cre mice RRID:MMRRC_030989-UCD, C57BL/6J A2A-Cre mice RRID:MMRRC_036158-UCD

Abstract

Balancing experimental control with ecological validity remains a central challenge for studying brain function. Here, we developed the Tower Foraging Park (TFP), a self-paced behavioral paradigm that emulates patch foraging and in which mice collect rewards by performing directional quarter-turns around square towers (exploit) and switching between them as they become depleted (explore). Mice rapidly learned the rewarded turning direction, increased movement speed, reduced trajectory variability, and typically abandoned towers before depletion. Reversal of the rewarded direction triggered rapid adaptation, accompanied by a dissociation between movement variability and speed. Repeated reversals progressively improved flexibility. Increasing the difficulty of locating rewarding towers prolonged exploitation, revealing adaptive regulation of patch-leaving decisions. Finally, mice trained under flexible contingencies ultimately outperformed those trained under stable contingencies in a challenging context. Altogether, the TFP reveals mechanisms underlying flexible foraging and provides a versatile, ecologically grounded platform for investigating the neural bases of adaptive behavior.

Reproduced under the paper's license (CC BY-NC), from the paper cited above.

Repositories

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

robbe-neuroteam/TowerForagingPark_Hardware

License: other
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: 4baa152f6bf123d3c39a4824c07eb9de15d8ee7c, 29 May 2026
Size: 70 files
Software Heritage: not archived
Found in: “Data and code availability”
Holds: README, license file, documentation
Not found: CITATION.cff, environment file, tests, continuous integration
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
2 files

robbe-neuroteam/TowerForagingPark_iScience_Figures

License: none: the authors keep all their rights
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: e3bdaa3bb781002c1457da8fce19b36545539420, 17 June 2026
Languages: Jupyter (38), Python (2)
Size: 112 files, 40 scripts
Software Heritage: not archived
Found in: “Data and code availability”
Holds: README, 38 notebooks
Not found: license file, CITATION.cff, environment file, tests, continuous integration, documentation
Tools: Matplotlib (13 files), NumPy (13 files), Pillow (12 files), SciPy (12 files)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
16 files

The paper's code and data availability statement is in the Data section.

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:

  • 2 repositories of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
  • 16 scripts, each with its path and the digest of its content;
  • 14 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

Data and code availability

• Group-wise raw and processed experimental data have been deposited at OSF and is available at https://osf.io/23f76/ in Group1Data and Group2Data as CSV (raw data) and PKL (processed data) files, and are classified by mouse and by session. They are publicly available as of the date of publication. Details on the experimental setup have been deposited at GitHub and are publicly available at https://github.com/robbe-neuroteam/TowerForagingPark_Hardware as of the date of publication. • All original analysis codes have been deposited at GitHub and are publicly available at https://github.com/robbe-neuroteam/TowerForagingPark_iScience_Figures as of the date of publication. • Any additional information required to reanalyze the data reported in this paper is available from the lead contact upon request.

Reproduced under the paper's license (CC BY-NC), from the paper cited above.

Materials availability

This study did not generate biological reagents or animal lines. Custom-built components of the TFP apparatus and custom codes used for behavioral acquisition and analysis are publicly available through the GitHub repository indicated in the data and code availability section.

Reproduced under the paper's license (CC BY-NC), from the paper cited above.

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, 27 September 2026: the first record

Recorded: type, language, journal, volume, issue, pages, dates, 7 authors, 3 keywords, 2 funders, 94 references, 3 RRIDs.

Cite

This paper

Schaffhauser, M., Orjollet-Lacomme, T., Amroune, K., Morvan, T., Fortoul, A., Lechelon, M., & Robbe, D. (2026). The Tower Foraging Park: A paradigm for studying cognitive and motor processes underlying behavioral flexibility in freely moving mice. iScience, 29(7), 116498. https://doi.org/10.1016/j.isci.2026.116498

BibTeX

@article{schaffhauser2026tower,
author = {Schaffhauser, Maud and Orjollet-Lacomme, Tom and Amroune, Kenza and Morvan, Thomas and Fortoul, Aurélien and Lechelon, Mathias and Robbe, David},
title = {{The Tower Foraging Park: A paradigm for studying cognitive and motor processes underlying behavioral flexibility in freely moving mice}},
journal = {iScience},
year = {2026},
month = jun,
volume = {29},
number = {7},
pages = {116498},
publisher = {Elsevier},
issn = {2589-0042},
doi = {10.1016/j.isci.2026.116498},
url = {https://doi.org/10.1016/j.isci.2026.116498},
pmid = {42491644},
pmcid = {PMC13377965}
}

RIS

TY - JOUR
AU - Schaffhauser, Maud
AU - Orjollet-Lacomme, Tom
AU - Amroune, Kenza
AU - Morvan, Thomas
AU - Fortoul, Aurélien
AU - Lechelon, Mathias
AU - Robbe, David
TI - The Tower Foraging Park: A paradigm for studying cognitive and motor processes underlying behavioral flexibility in freely moving mice
T2 - iScience
J2 - iScience
PY - 2026
DA - 2026/06/26
VL - 29
IS - 7
SP - 116498
SN - 2589-0042
PB - Elsevier
DO - 10.1016/j.isci.2026.116498
UR - https://doi.org/10.1016/j.isci.2026.116498
LA - en
ER -

CSL-JSON

{
"id": "10.1016/j.isci.2026.116498",
"type": "article-journal",
"title": "The Tower Foraging Park: A paradigm for studying cognitive and motor processes underlying behavioral flexibility in freely moving mice",
"container-title": "iScience",
"author": [
{
"family": "Schaffhauser",
"given": "Maud"
},
{
"family": "Orjollet-Lacomme",
"given": "Tom"
},
{
"family": "Amroune",
"given": "Kenza"
},
{
"family": "Morvan",
"given": "Thomas"
},
{
"family": "Fortoul",
"given": "Aurélien"
},
{
"family": "Lechelon",
"given": "Mathias"
},
{
"family": "Robbe",
"given": "David"
}
],
"container-title-short": "iScience",
"volume": "29",
"issue": "7",
"page": "116498",
"DOI": "10.1016/j.isci.2026.116498",
"PMID": "42491644",
"PMCID": "PMC13377965",
"ISSN": "2589-0042",
"publisher": "Elsevier",
"URL": "https://doi.org/10.1016/j.isci.2026.116498",
"language": "en",
"issued": {
"date-parts": [
[
2026,
6,
26
]
]
}
}

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.1038/s41467-026-73994-1 [code]
Prediction error correlates in the striosome-dopamine circuit emerge from information gain.
Journal: Nature communications
In common: SciPy, Matplotlib, NumPy, 6 references
[2] doi:10.1038/s41467-026-72057-9 [code]
Sex-specific behavioral feedback modulates sensorimotor processing and drives flexible social behavior.
Journal: Nature communications
In common: Pillow, SciPy, Matplotlib, 1 other tool, 3 references
[3] doi:10.1038/s41586-026-10528-1 [code]
A critical initialization for biological neural networks.
Journal: Nature
In common: SciPy, Matplotlib, NumPy, mouse, 4 references
[4] doi:10.1038/s41467-026-74227-1 [code]
Age-related changes in behavioural and neural variability in a decision-making task.
Journal: Nature communications
In common: SciPy, Matplotlib, NumPy, mouse, 4 references
[5] doi:10.7554/elife.111876 [code]
Distinct sensorimotor encoding in tuft dendrites and somata associated with action, correction, and learning.
Journal: eLife
In common: Pillow, SciPy, Matplotlib, 1 other tool, mouse, 3 references
[6] doi:10.7554/elife.109717 [code]
Retrosplenial cortex enables context-dependent goal-directed sensorimotor transformation.
Journal: eLife
In common: Pillow, SciPy, Matplotlib, 1 other tool, mouse, 3 references
[7] doi:10.1126/sciadv.aef0343 [code]
Learning induces activation-mechanism-dependent neural plasticity in an intracortical microstimulation task.
Journal: Science advances
In common: Pillow, SciPy, Matplotlib, 1 other tool, 3 references
[8] doi:10.7554/elife.109571 [code]
Rank- and threat-dependent social modulation of innate defensive behaviors.
Journal: eLife
In common: SciPy, Matplotlib, NumPy, mouse, 3 references
[9] 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: SciPy, Matplotlib, NumPy, mouse, 3 references
[10] doi:10.1016/j.crmeth.2026.101421 [code]
EthoPy provides an accessible platform for reproducible behavioral neuroscience.
Journal: Cell reports methods
In common: SciPy, Matplotlib, NumPy, mouse, 3 references

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.