OSCR

Rank- and threat-dependent social modulation of innate defensive behaviors.

Code ↔ Paper

The paper beside its authors' code: matches between them have not been computed for this paper yet.

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 · 5,049 lines · 195 KB · Apache-2.0

  1. # %%
  2. # Reload modules automatically.
  3. %load_ext autoreload
  4. %autoreload 2
  5. import os
  6. print(os.getcwd())
  7. import pandas as pd
  8. import matplotlib.pyplot as plt
  9. import seaborn as sns
  10. import numpy as np
  11. # %% [markdown]
  12. # ## Load manual labeled data and animal tracking data.
  13. # %% [markdown]
  14. # ### Load lst coordinate data and behavior data.
  15. # %%
  16. import os
  17. import pickle
  18. from analyze_data_utils import lst_read_deepethogram_csv, lst_read_bento_annot, lst_read_decision_csv, lst_read_deeplabcut_h5_and_analyze_data
  19. # lst_decision_file = os.path.abspath(os.path.join('.', 'lst', 'looming_behavior_sb_1looming_lessn.csv'))
  20. # lst_deci_dict = lst_read_decision_csv(lst_decision_file)
  21. # with open('lst_deci_dict.pkl', 'wb') as f:
  22. # pickle.dump(lst_deci_dict, f)
  23. with open('lst_deci_dict.pkl', 'rb') as f:
  24. lst_deci_dict = pickle.load(f)
  25. # lst_labels_dir = os.path.abspath(os.path.join('.', 'lst', 'deg_csv_data'))
  26. # lst_bento_corr_file = os.path.abspath(os.path.join('.', 'lst', 'bento_correction_sb_1looming.csv'))
  27. # lst_bhvr_frames_dict, lst_baseline_bhvr_frames_dict = lst_read_deepethogram_csv(lst_labels_dir, lst_bento_corr_file)
  28. # with open('lst_bhvr_frames_dict.pkl', 'wb') as f:
  29. # pickle.dump(lst_bhvr_frames_dict, f)
  30. # with open('lst_baseline_bhvr_frames_dict.pkl', 'wb') as f:
  31. # pickle.dump(lst_baseline_bhvr_frames_dict, f)
  32. with open('lst_bhvr_frames_dict.pkl', 'rb') as f:
  33. lst_bhvr_frames_dict = pickle.load(f)
  34. with open('lst_baseline_bhvr_frames_dict.pkl', 'rb') as f:
  35. lst_baseline_bhvr_frames_dict = pickle.load(f)
  36. # lst_bento_dir = os.path.abspath(os.path.join('.', 'lst', 'bento_annot_data'))
  37. # lst_bhvr_labels_dict, lst_baseline_bhvr_labels_dict = lst_read_bento_annot(lst_bento_dir, lst_bento_corr_file, lst_deci_dict)
  38. # with open('lst_bhvr_labels_dict.pkl', 'wb') as f:
  39. # pickle.dump(lst_bhvr_labels_dict, f)
  40. # with open('lst_baseline_bhvr_labels_dict.pkl', 'wb') as f:
  41. # pickle.dump(lst_baseline_bhvr_labels_dict, f)
  42. with open('lst_bhvr_labels_dict.pkl', 'rb') as f:
  43. lst_bhvr_labels_dict = pickle.load(f)
  44. with open('lst_baseline_bhvr_labels_dict.pkl', 'rb') as f:
  45. lst_baseline_bhvr_labels_dict = pickle.load(f)
  46. # lst_dlc_h5_dir = os.path.abspath(os.path.join('.', 'lst', 'dlc_h5_data'))
  47. # lsts_data, lstp_data = lst_read_deeplabcut_h5_and_analyze_data(lst_dlc_h5_dir)
  48. # with open('lsts_data.pkl', 'wb') as f:
  49. # pickle.dump(lsts_data, f)
  50. # with open('lstp_data.pkl', 'wb') as f:
  51. # pickle.dump(lstp_data, f)
  52. with open('lsts_data.pkl', 'rb') as f:
  53. lsts_data = pickle.load(f)
  54. with open('lstp_data.pkl', 'rb') as f:
  55. lstp_data = pickle.load(f)
  56. # %% [markdown]
  57. # ### Global variables: ms_ids, mp_ids, looming_start_frames
  58. # %%
  59. ms_ids = ['B1M1(S)', 'B1M2(D)', 'B2M1(D)', 'B2M2(S)', 'B3M1(S)', 'B3M2(D)', 'B4M1(D)', 'B4M2(S)',
  60. 'B5M1(D)', 'B5M2(S)', 'B6M1(D)', 'B6M2(S)', 'B7M1(S)', 'B7M2(D)', 'B8M1(D)', 'B8M2(S)',
  61. 'B9M1(D)', 'B9M2(S)', 'B10M1(S)', 'B10M2(D)', 'B11M1(S)', 'B11M2(D)', 'B12M1(S)', 'B12M2(D)',
  62. 'B13M1(S)', 'B13M2(D)', 'B14M1(S)', 'B14M2(D)', 'B15M1(S)', 'B15M2(D)', 'B16M1(S)', 'B16M2(D)',
  63. 'B17M1(S)', 'B17M2(D)', 'B18M1(D)', 'B18M2(S)', 'B19M1(D)', 'B19M2(S)', 'B20M1(S)', 'B20M2(D)']
  64. mp_ids = ['B1M1(S)_B1M2(D)', 'B2M1(D)_B2M2(S)', 'B3M1(S)_B3M2(D)', 'B4M1(D)_B4M2(S)',
  65. 'B5M1(D)_B5M2(S)', 'B6M1(D)_B6M2(S)', 'B7M1(S)_B7M2(D)', 'B8M1(D)_B8M2(S)',
  66. 'B9M1(D)_B9M2(S)', 'B10M1(S)_B10M2(D)', 'B11M1(S)_B11M2(D)', 'B12M1(S)_B12M2(D)',
  67. 'B13M1(S)_B13M2(D)', 'B14M1(S)_B14M2(D)', 'B15M1(S)_B15M2(D)', 'B16M1(S)_B16M2(D)',
  68. 'B17M1(S)_B17M2(D)', 'B18M1(D)_B18M2(S)', 'B19M1(D)_B19M2(S)', 'B20M1(S)_B20M2(D)']
  69. mS_ids = ['B1M1(S)', 'B2M2(S)', 'B3M1(S)', 'B4M2(S)', 'B5M2(S)', 'B6M2(S)', 'B7M1(S)', 'B8M2(S)',
  70. 'B9M2(S)', 'B10M1(S)', 'B11M1(S)', 'B12M1(S)', 'B13M1(S)', 'B14M1(S)', 'B15M1(S)', 'B16M1(S)',
  71. 'B17M1(S)', 'B18M2(S)', 'B19M2(S)', 'B20M1(S)']
  72. mD_ids = ['B1M2(D)', 'B2M1(D)', 'B3M2(D)', 'B4M1(D)', 'B5M1(D)', 'B6M1(D)', 'B7M2(D)', 'B8M1(D)',
  73. 'B9M1(D)', 'B10M2(D)', 'B11M2(D)', 'B12M2(D)', 'B13M2(D)', 'B14M2(D)', 'B15M2(D)', 'B16M2(D)',
  74. 'B17M2(D)', 'B18M1(D)', 'B19M1(D)', 'B20M2(D)']
  75. looming_start_frames = {'B1M1(S)': [21687, 24882],
  76. 'B1M2(D)': [19787, 27340],
  77. 'B2M1(D)': [20822, 28112],
  78. 'B2M2(S)': [20517, 27347],
  79. 'B3M1(S)': [23645, 27281],
  80. 'B3M2(D)': [18638, 26129],
  81. 'B4M1(D)': [19250, 25754],
  82. 'B4M2(S)': [19919, 25135],
  83. 'B5M1(D)': [18124, 25424],
  84. 'B5M2(S)': [18522, 33582],
  85. 'B6M1(D)': [18365, 25547, 29965],
  86. 'B6M2(S)': [18517, 22786, 27173, 34537],
  87. 'B7M1(S)': [20516, 31677, 45430, 55860, 67593],
  88. 'B7M2(D)': [18519, 30751, 35634, 40217, 43675],
  89. 'B8M1(D)': [19927, 22999, 32861, 39477, 47451],
  90. 'B8M2(S)': [20228, 26561, 34729, 36654, 45224],
  91. 'B9M1(D)': [19173, 21997, 25862, 28980, 32281],
  92. 'B9M2(S)': [20031, 24378, 28080, 32859, 37854],
  93. 'B10M1(S)': [18291, 22143, 25473, 31536, 43488],
  94. 'B10M2(D)': [19743, 24225, 29010, 34332, 38145],
  95. 'B11M1(S)': [18180, 22365, 26364, 29664, 32778],
  96. 'B11M2(D)': [22305, 29619, 35238, 42942, 52587],
  97. 'B12M1(S)': [26184, 35190, 38490, 44826, 63663],
  98. 'B12M2(D)': [22101, 26424, 28731, 33624, 38262],
  99. 'B13M1(S)': [19734, 21690, 30522, 33570, 39768],
  100. 'B13M2(D)': [19563, 22602, 36426, 43215, 49560],
  101. 'B14M1(S)': [21639, 24939, 29826, 33582, 38907],
  102. 'B14M2(D)': [25443, 28803, 32490, 42564, 52275],
  103. 'B15M1(S)': [18798, 24028, 27337, 32670, 39426, 45213, 50346, 53076],
  104. 'B15M2(D)': [18759, 25234, 30829, 34297, 37946, 43971, 48324, 51170, 64290, 69035],
  105. 'B16M1(S)': [18354, 20851, 25998, 31204, 35160, 40137, 51816],
  106. 'B16M2(D)': [18282, 25614, 28758, 33453, 35790, 41724, 52158, 69162],
  107. 'B17M1(S)': [18922, 22083, 25744, 33987, 37318, 44458, 53217, 61588, 65379],
  108. 'B17M2(D)': [18907, 23739, 30816, 41705, 53482],
  109. 'B18M1(D)': [18645, 21592, 45597, 52743],
  110. 'B18M2(S)': [18899, 21992, 28926, 36918, 40850, 54897, 58513, 61513],
  111. 'B19M1(D)': [19371, 27168, 42195, 50253, 58077],
  112. 'B19M2(S)': [19974, 28379],
  113. 'B20M1(S)': [22228, 26274, 55110, 60789, 71464],
  114. 'B20M2(D)': [18991, 24878, 28582, 35260, 38830, 45945, 51220, 55177, 60471, 64398],
  115. 'B1M1(S)_B1M2(D)': [26897, 32648],
  116. 'B2M1(D)_B2M2(S)': [53137, 60549, 70481],
  117. 'B3M1(S)_B3M2(D)': [29541, 37526, 40262, 47454, 69994],
  118. 'B5M1(D)_B5M2(S)': [25255, 28961, 39435, 43291, 70922],
  119. 'B4M1(D)_B4M2(S)': [20228, 23426, 26259, 28705, 33779],
  120. 'B6M1(D)_B6M2(S)': [25041, 29475, 40256, 45849, 63113],
  121. 'B7M1(S)_B7M2(D)': [22068, 33550, 52853, 57657, 65416],
  122. 'B8M1(D)_B8M2(S)': [22155, 31191, 34569, 51364, 70682],
  123. 'B9M1(D)_B9M2(S)': [18831, 20865, 24561, 28755, 33870, 45141, 49365, 52656, 69148],
  124. 'B10M1(S)_B10M2(D)': [22422, 27787, 36190, 41217, 45228, 48732, 52455, 56308, 60924, 66468],
  125. 'B11M1(S)_B11M2(D)': [19539, 23068, 26653, 30262, 34231, 37920, 41037, 46828, 49966, 53490],
  126. 'B12M1(S)_B12M2(D)': [19257, 23934, 28201, 30547, 34242, 37833, 42060, 47703, 52767, 57060],
  127. 'B13M1(S)_B13M2(D)': [19183, 23344, 29673, 33379, 35610, 41544, 48450, 51997, 56022, 63910],
  128. 'B14M1(S)_B14M2(D)': [22887, 29199, 40301, 48266, 52510, 67180, 71367],
  129. 'B15M1(S)_B15M2(D)': [18831, 22671, 28074, 31942, 36654, 39847, 44067, 49578, 53272, 57154],
  130. 'B16M1(S)_B16M2(D)': [20454, 29139, 32550, 36150, 40696, 44059, 49290, 52071, 58332, 63426],
  131. 'B17M1(S)_B17M2(D)': [25080, 33249, 37938, 45118, 51993, 55266, 60331, 63778, 68193, 71500],
  132. 'B18M1(D)_B18M2(S)': [18876, 24573, 28573, 32257, 36224, 38827, 42231, 44521, 49630, 56858],
  133. 'B19M1(D)_B19M2(S)': [28854, 33918, 37251, 41328, 46119, 49065, 51954, 56518, 61744, 64632],
  134. 'B20M1(S)_B20M2(D)': [22090, 27171, 31389, 37048, 40728, 45769, 48780, 52728, 54745, 58203]}
  135. lst_pixpercm = 21.76
  136. framerate = 30
  137. # %%
  138. lst_all_bhvrs = [
  139. 'escape',
  140. 'freezing',
  141. 'tail_rattling',
  142. 'rearing/up_stretch',
  143. 'huddling',
  144. 'approach_partner',
  145. 'follow_partner',
  146. 'groom_partner',
  147. 'sniff_partner',
  148. 'grooming',
  149. 'sniffing',
  150. 'gap_state',
  151. ]
  152. lst_bhvr_abbrev_dict = {
  153. "escape": 'E',
  154. "freezing": 'F',
  155. "tail_rattling": 'T',
  156. "rearing/up_stretch": 'R',
  157. "huddling": 'H',
  158. "approach_partner": 'AP',
  159. "follow_partner": 'FP',
  160. "groom_partner": 'GP',
  161. "sniff_partner": 'SP',
  162. "grooming": 'G',
  163. "sniffing": 'S',
  164. 'gap_state': 'O'
  165. }
  166. lst_bhvr_color_dict = {
  167. "E": "#FF1717",
  168. "R": "#C409FF",
  169. "F": "#FF75EF",
  170. "T": "#FF09FF",
  171. "H": "#6D57F3",
  172. "AP": "#143FCA",
  173. "FP": "#137CAB",
  174. "GP": "#0990FF",
  175. "SP": "#09FFFF",
  176. "G": "#2EA761",
  177. "S": "#0FD400",
  178. "O": "white",
  179. 'Def': 'red',
  180. 'Soc': 'royalblue',
  181. 'Exp&\nRest': 'limegreen'
  182. }
  183. lst_arena_params={
  184. 'pixpergrid': 43.52,
  185. 'arena_size': 1088,
  186. 'edge_width': 43.52 * 4,
  187. 'nest_coord': [0, 43.52 * 17.5],
  188. 'nest_size': [43.52 * 10, 43.52 * 7.5]
  189. }
  190. # %% [markdown]
  191. # ### Move csv files and convert csv to annot.
  192. # %%
  193. import os
  194. from move_csv_files import copy_csv_files_to_root, move_csv_files_to_target_folders
  195. lsts_deepethogram_project_DATA_dir = os.path.abspath(os.path.join('.', 'lst', 'lsts_deepethogram', 'DATA'))
  196. copy_csv_files_to_root(lsts_deepethogram_project_DATA_dir)
  197. lstp_deepethogram_project_DATA_dir = os.path.abspath(os.path.join('.', 'lst', 'lstp_deepethogram', 'DATA'))
  198. copy_csv_files_to_root(lstp_deepethogram_project_DATA_dir)
  199. lst_labels_dir = os.path.abspath(os.path.join('.', 'lst', 'deg_csv_data'))
  200. temp_lsts_labels_dir = os.path.abspath(os.path.join('.', 'lst', 'labels_csv_lsts'))
  201. temp_lstp_labels_dir = os.path.abspath(os.path.join('.', 'lst', 'labels_csv_lstp'))
  202. move_csv_files_to_target_folders(lsts_deepethogram_project_DATA_dir, temp_lsts_labels_dir)
  203. move_csv_files_to_target_folders(lstp_deepethogram_project_DATA_dir, temp_lstp_labels_dir)
  204. # %%
  205. from convert_labelscsv2annot_lst import get_csv_files, read_labelscsv_file, create_annot_content, creat_annot_file
  206. lst_bento_dir = os.path.abspath(os.path.join('.', 'lst', 'bento_annot_data'))
  207. group = 's'
  208. csv_files = get_csv_files(temp_lsts_labels_dir)
  209. for csv_file in csv_files:
  210. csv_file_path = os.path.join(temp_lsts_labels_dir, csv_file)
  211. behavior_frames_dict = read_labelscsv_file(csv_file_path, group)
  212. annot_content = create_annot_content(csv_file, group)
  213. creat_annot_file(csv_file, lst_bento_dir, annot_content, behavior_frames_dict)
  214. group = 'p'
  215. csv_files = get_csv_files(temp_lstp_labels_dir)
  216. for csv_file in csv_files:
  217. csv_file_path = os.path.join(temp_lstp_labels_dir, csv_file)
  218. behavior_frames_dict = read_labelscsv_file(csv_file_path, group)
  219. annot_content = create_annot_content(csv_file, group)
  220. creat_annot_file(csv_file, lst_bento_dir, annot_content, behavior_frames_dict)
  221. # %%
  222. move_csv_files_to_target_folders(temp_lsts_labels_dir, lst_labels_dir)
  223. move_csv_files_to_target_folders(temp_lstp_labels_dir, lst_labels_dir)
  224. # %% [markdown]
  225. # ### Check lst overlap.
  226. # %%
  227. from collections import Counter
  228. # Set frame range
  229. frame_start = 150
  230. frame_end = 1950
  231. total_frames = frame_end - frame_start
  232. # Excluded behaviors
  233. excluded_behaviors = ['nest']
  234. # Iterate through all groups and ranks
  235. for group in ['s', 'p']:
  236. group_name = 'Single' if group == 's' else 'Pair'
  237. print(f"\n{'='*60}")
  238. print(f"Group: {group_name}")
  239. print('='*60)
  240. for rank in ['D', 'S']:
  241. rank_name = 'Dominant' if rank == 'D' else 'Subordinate'
  242. print(f"\n--- Rank: {rank_name} ---")
  243. if group not in lst_bhvr_frames_dict or rank not in lst_bhvr_frames_dict[group]:
  244. print(f" No data for {group_name} {rank_name}")
  245. continue
  246. # Initialize aggregated counter for all trials in this group-rank combination
  247. aggregated_combinations = []
  248. aggregated_overlap_frames = [] # Store all overlapping frames across trials
  249. trial_count = 0
  250. # Iterate through all trials
  251. for t_id in sorted(lst_bhvr_frames_dict[group][rank].keys()):
  252. trial_data = lst_bhvr_frames_dict[group][rank][t_id]
  253. # Get all available behaviors for this trial
  254. available_behaviors = [b for b in trial_data.keys() if b not in excluded_behaviors]
  255. if len(available_behaviors) == 0:
  256. print(f"\n {t_id}: No behaviors available after exclusion")
  257. continue
  258. # Create a list to store behavior combination for each frame
  259. frame_combinations = []
  260. overlap_frames_info = [] # Store frame numbers with overlapping behaviors
  261. # Iterate through each frame
  262. for frame_idx in range(frame_start, frame_end):
  263. # Find all active behaviors at this frame
  264. active_behaviors = []
  265. for behavior in sorted(available_behaviors): # Sort for consistent ordering
  266. try:
  267. if frame_idx in trial_data[behavior].index and trial_data[behavior][frame_idx] == 1:
  268. active_behaviors.append(behavior)
  269. except:
  270. pass
  271. # priority_behaviors = ['stretching']
  272. # if 'sniffing' in active_behaviors:
  273. # for priority_bhvr in priority_behaviors:
  274. # if priority_bhvr in active_behaviors:
  275. # active_behaviors.remove('sniffing')
  276. # break
  277. # Create combination string
  278. if len(active_behaviors) == 0:
  279. combo_str = "no_behavior"
  280. elif len(active_behaviors) == 1:
  281. combo_str = active_behaviors[0]
  282. else:
  283. combo_str = "+".join(active_behaviors)
  284. # Record overlapping frame information
  285. overlap_frames_info.append({
  286. 'frame': frame_idx,
  287. 'behaviors': active_behaviors
  288. })
  289. frame_combinations.append(combo_str)
  290. # Add this trial's combinations to the aggregated list
  291. aggregated_combinations.extend(frame_combinations)
  292. # Add overlapping frames info with trial ID
  293. for frame_info in overlap_frames_info:
  294. aggregated_overlap_frames.append({
  295. 'trial': t_id,
  296. 'frame': frame_info['frame'],
  297. 'behaviors': frame_info['behaviors']
  298. })
  299. trial_count += 1
  300. # Count occurrences of each combination
  301. combination_counter = Counter(frame_combinations)
  302. # Calculate percentages and sort by frequency
  303. combination_stats = []
  304. for combo, count in combination_counter.items():
  305. pct = (count / total_frames) * 100
  306. combination_stats.append({
  307. 'combination': combo,
  308. 'frames': count,
  309. 'pct': pct
  310. })
  311. # Sort by frame count (descending)
  312. combination_stats.sort(key=lambda x: x['frames'], reverse=True)
  313. # Print results
  314. print(f"\n {t_id}:")
  315. print(f" Total unique behavior combinations: {len(combination_stats)}")
  316. print(f" Frame distribution:")
  317. # Show all combinations
  318. for stat in combination_stats:
  319. print(f" {stat['combination']}: {stat['frames']} frames ({stat['pct']:.2f}%)")
  320. # Display overlapping frames details
  321. if len(overlap_frames_info) > 0:
  322. print(f"\n Overlapping frames details ({len(overlap_frames_info)} frames total):")
  323. # Show first 20 overlapping frames as examples
  324. max_display = min(20, len(overlap_frames_info))
  325. for i, frame_info in enumerate(overlap_frames_info[:max_display]):
  326. behaviors_str = " + ".join(frame_info['behaviors'])
  327. print(f" Frame {frame_info['frame']}: {behaviors_str}")
  328. if len(overlap_frames_info) > max_display:
  329. print(f" ... and {len(overlap_frames_info) - max_display} more overlapping frames")
  330. # Summary: single vs multiple behaviors
  331. single_behavior_frames = sum(stat['frames'] for stat in combination_stats
  332. if '+' not in stat['combination'] and stat['combination'] != 'no_behavior')
  333. multiple_behavior_frames = sum(stat['frames'] for stat in combination_stats
  334. if '+' in stat['combination'])
  335. no_behavior_frames = sum(stat['frames'] for stat in combination_stats
  336. if stat['combination'] == 'no_behavior')
  337. print(f"\n Summary:")
  338. print(f" Single behavior: {single_behavior_frames} frames ({single_behavior_frames/total_frames*100:.2f}%)")
  339. print(f" Multiple behaviors: {multiple_behavior_frames} frames ({multiple_behavior_frames/total_frames*100:.2f}%)")
  340. print(f" No behavior: {no_behavior_frames} frames ({no_behavior_frames/total_frames*100:.2f}%)")
  341. # Print aggregated statistics for this group-rank combination
  342. if trial_count > 0:
  343. print(f"\n{'*' * 50}")
  344. print(f"AGGREGATED STATISTICS: {group_name} - {rank_name}")
  345. print(f"Total trials: {trial_count}")
  346. print(f"Total frames analyzed: {len(aggregated_combinations)}")
  347. print(f"{'*' * 50}")
  348. # Count aggregated combinations
  349. aggregated_counter = Counter(aggregated_combinations)
  350. total_aggregated_frames = len(aggregated_combinations)
  351. # Calculate percentages and sort
  352. aggregated_stats = []
  353. for combo, count in aggregated_counter.items():
  354. pct = (count / total_aggregated_frames) * 100
  355. aggregated_stats.append({
  356. 'combination': combo,
  357. 'frames': count,
  358. 'pct': pct
  359. })
  360. aggregated_stats.sort(key=lambda x: x['frames'], reverse=True)
  361. print(f"\nAggregated behavior combinations:")
  362. for stat in aggregated_stats:
  363. print(f" {stat['combination']}: {stat['frames']} frames ({stat['pct']:.2f}%)")
  364. # Display all overlapping frames in aggregated data
  365. if len(aggregated_overlap_frames) > 0:
  366. print(f"\n All overlapping frames across trials ({len(aggregated_overlap_frames)} frames total):")
  367. # Group by trial for better readability
  368. current_trial = None
  369. max_display_per_trial = 10
  370. trial_frame_count = 0
  371. for frame_info in aggregated_overlap_frames:
  372. if current_trial != frame_info['trial']:
  373. if current_trial is not None and trial_frame_count > max_display_per_trial:
  374. print(f" ... and {trial_frame_count - max_display_per_trial} more frames in this trial")
  375. current_trial = frame_info['trial']
  376. trial_frame_count = 0
  377. print(f"\n {current_trial}:")
  378. trial_frame_count += 1
  379. if trial_frame_count <= max_display_per_trial:
  380. behaviors_str = " + ".join(frame_info['behaviors'])
  381. print(f" Frame {frame_info['frame']}: {behaviors_str}")
  382. if trial_frame_count > max_display_per_trial:
  383. print(f" ... and {trial_frame_count - max_display_per_trial} more frames in this trial")
  384. # Aggregated summary
  385. agg_single = sum(stat['frames'] for stat in aggregated_stats
  386. if '+' not in stat['combination'] and stat['combination'] != 'no_behavior')
  387. agg_multiple = sum(stat['frames'] for stat in aggregated_stats
  388. if '+' in stat['combination'])
  389. agg_no_behavior = sum(stat['frames'] for stat in aggregated_stats
  390. if stat['combination'] == 'no_behavior')
  391. print(f"\nAggregated Summary:")
  392. print(f" Single behavior: {agg_single} frames ({agg_single/total_aggregated_frames*100:.2f}%)")
  393. print(f" Multiple behaviors: {agg_multiple} frames ({agg_multiple/total_aggregated_frames*100:.2f}%)")
  394. print(f" No behavior: {agg_no_behavior} frames ({agg_no_behavior/total_aggregated_frames*100:.2f}%)")
  395. print(f"{'*' * 50}\n")
  396. print("\n" + "="*80)
  397. print("Frame-by-frame analysis complete")
  398. print("="*80)
  399. # %% [markdown]
  400. # ### Load ret coordinate data and behavior data.
  401. # %%
  402. import os
  403. import pickle
  404. from analyze_data_utils import ret_read_deepethogram_csv, ret_read_bento_annot, ret_read_ethovision_xlsx_and_analyze_data, reorder_dict_sessions
  405. session_order = ['B1', 'B2', 'B3', 'B4', 'B5', 'B6', 'B7', 'B8', 'B9', 'B10', 'B11']
  406. # ret_labels_dir = os.path.abspath(os.path.join('.', 'ret', 'deg_csv_data'))
  407. # ret_bento_corr_file = os.path.abspath(os.path.join('.', 'ret', 'bento_correction_sb_ret.csv'))
  408. # ret_bhvr_frames_dict, ret_origin_bhvr_frames_dict = ret_read_deepethogram_csv(ret_labels_dir, ret_bento_corr_file)
  409. # ret_bhvr_frames_dict = reorder_dict_sessions(ret_bhvr_frames_dict, session_order)
  410. # ret_origin_bhvr_frames_dict = reorder_dict_sessions(ret_origin_bhvr_frames_dict, session_order)
  411. # with open('ret_bhvr_frames_dict.pkl', 'wb') as f:
  412. # pickle.dump(ret_bhvr_frames_dict, f)
  413. # with open('ret_origin_bhvr_frames_dict.pkl', 'wb') as f:
  414. # pickle.dump(ret_origin_bhvr_frames_dict, f)
  415. with open('ret_bhvr_frames_dict.pkl', 'rb') as f:
  416. ret_bhvr_frames_dict = pickle.load(f)
  417. with open('ret_origin_bhvr_frames_dict.pkl', 'rb') as f:
  418. ret_origin_bhvr_frames_dict = pickle.load(f)
  419. # ret_bento_dir = os.path.abspath(os.path.join('.', 'ret', 'bento_annot_data'))
  420. # ret_bhvr_labels_dict, ret_origin_bhvr_labels_dict = ret_read_bento_annot(ret_bento_dir, ret_bento_corr_file)
  421. # ret_bhvr_labels_dict = reorder_dict_sessions(ret_bhvr_labels_dict, session_order)
  422. # ret_origin_bhvr_labels_dict = reorder_dict_sessions(ret_origin_bhvr_labels_dict, session_order)
  423. # with open('ret_bhvr_labels_dict.pkl', 'wb') as f:
  424. # pickle.dump(ret_bhvr_labels_dict, f)
  425. # with open('ret_origin_bhvr_labels_dict.pkl', 'wb') as f:
  426. # pickle.dump(ret_origin_bhvr_labels_dict, f)
  427. with open('ret_bhvr_labels_dict.pkl', 'rb') as f:
  428. ret_bhvr_labels_dict = pickle.load(f)
  429. with open('ret_origin_bhvr_labels_dict.pkl', 'rb') as f:
  430. ret_origin_bhvr_labels_dict = pickle.load(f)
  431. # ret_ev_xlsx_dir = os.path.abspath(os.path.join('.', 'ret', 'ev_xlsx_data'))
  432. # ret_raw_data, ret_stat_data = ret_read_ethovision_xlsx_and_analyze_data(ret_ev_xlsx_dir)
  433. # with open('ret_raw_data.pkl', 'wb') as f:
  434. # pickle.dump(ret_raw_data, f)
  435. # with open('ret_stat_data.pkl', 'wb') as f:
  436. # pickle.dump(ret_stat_data, f)
  437. with open('ret_raw_data.pkl', 'rb') as f:
  438. ret_raw_data = pickle.load(f)
  439. with open('ret_stat_data.pkl', 'rb') as f:
  440. ret_stat_data = pickle.load(f)
  441. # %% [markdown]
  442. # ### Global variables: rat_in_frames
  443. # %%
  444. rat_in_frames = {
  445. 'D1M1&M2': 11100,
  446. 'D1M1': 10766,
  447. 'D1M2': 10713,
  448. 'D2M1&M2': 10416,
  449. 'D2M1': 10651,
  450. 'D2M2': 10633,
  451. 'D3M1&M2': 10462,
  452. 'D3M1': 10987,
  453. 'D3M2': 10677,
  454. 'D4M1&M2': 12051,
  455. 'D4M1': 11359,
  456. 'D4M2': 10424,
  457. 'D5M1&M2': 11100,
  458. 'D5M1': 10414,
  459. 'D5M2': 10395,
  460. 'D6M1&M2': 10751,
  461. 'D6M1': 10301,
  462. 'D6M2': 10420,
  463. 'D7M1&M2': 10374,
  464. 'D7M1': 10141,
  465. 'D7M2': 9586,
  466. 'D8M1&M2': 10082,
  467. 'D8M1': 10098,
  468. 'D8M2': 10197,
  469. 'D9M1&M2': 10188,
  470. 'D9M1': 10283,
  471. 'D9M2': 10441,
  472. 'D10M1&M2': 10510,
  473. 'D10M1': 10154,
  474. 'D10M2': 10343,
  475. 'D11M1&M2': 10463,
  476. 'D11M1': 10176,
  477. 'D11M2': 10413
  478. }
  479. framerate = 30
  480. # %%
  481. offset_frames = {
  482. 'D1M1&M2': 10812,
  483. 'D1M1': 10766,
  484. 'D1M2': 10713,
  485. 'D2M1&M2': 10286,
  486. 'D2M1': 10651,
  487. 'D2M2': 10633,
  488. 'D3M1&M2': 10470,
  489. 'D3M1': 10987,
  490. 'D3M2': 10677,
  491. 'D4M1&M2': 11817,
  492. 'D4M1': 11359,
  493. 'D4M2': 10424,
  494. 'D5M1&M2': 10976,
  495. 'D5M1': 10414,
  496. 'D5M2': 10395,
  497. 'D6M1&M2': 10624,
  498. 'D6M1': 10301,
  499. 'D6M2': 10420,
  500. 'D7M1&M2': 10225,
  501. 'D7M1': 10141,
  502. 'D7M2': 9586,
  503. 'D8M1&M2': 10066,
  504. 'D8M1': 10098,
  505. 'D8M2': 10197,
  506. 'D9M1&M2': 10196,
  507. 'D9M1': 10283,
  508. 'D9M2': 10441,
  509. 'D10M1&M2': 10556,
  510. 'D10M1': 10154,
  511. 'D10M2': 10343,
  512. 'D11M1&M2': 10942,
  513. 'D11M1': 10176,
  514. 'D11M2': 10413
  515. }
  516. # %%
  517. ret_all_bhvrs = [
  518. "approach",
  519. "investigation",
  520. "withdrawal",
  521. "stretch-attend",
  522. "freezing",
  523. "tail_rattling",
  524. "huddling",
  525. "approach_partner",
  526. "follow_partner",
  527. "groom_partner",
  528. "sniff_partner",
  529. "grooming",
  530. "gap_state",
  531. ]
  532. ret_bhvr_color_dict = {
  533. "A": "#FFCA09",
  534. "A+AP": "#8A856A",
  535. "I": "#FF8409",
  536. "W": "#FF0990",
  537. "W+AP": "#8A24AD",
  538. "S": "#FF7979",
  539. "F": "#FF75EF",
  540. "T": "#FF09FF",
  541. "H": "#6D57F3",
  542. "H+F": "#B666F1",
  543. "AP": "#143FCA",
  544. "FP": "#137CAB",
  545. "GP": "#0990FF",
  546. "SP": "#09FFFF",
  547. "G": "#28AE61",
  548. "O": "white",
  549. 'Def': 'red',
  550. 'Soc': 'royalblue',
  551. 'Exp&\nRest': 'limegreen'
  552. }
  553. ret_bhvr_abbrev_dict = {
  554. "approach": "A",
  555. "investigation": "I",
  556. "withdrawal": "W",
  557. "stretch-attend": "S",
  558. "freezing": "F",
  559. "tail_rattling": "T",
  560. "huddling": "H",
  561. "approach_partner": "AP",
  562. "follow_partner": "FP",
  563. "groom_partner": "GP",
  564. "sniff_partner": "SP",
  565. "grooming": "G",
  566. 'gap_state': 'O'
  567. }
  568. ret_arena_params={
  569. 'dx': 24/24, # pixel per cm
  570. 'rat_x': -5,
  571. 'near_width': 9,
  572. 'middle_width': 9,
  573. 'far_width': 6,
  574. 'arena_size': 24 * (24/24)
  575. }
  576. # %% [markdown]
  577. # ### Move csv files and convert csv to annot.
  578. # %%
  579. import os
  580. from move_csv_files import copy_csv_files_to_root, move_csv_files_to_target_folders, rename_csv_files_remove_video_prefix
  581. rets_deepethogram_project_DATA_dir = os.path.abspath(os.path.join('.', 'ret', 'rets_deepethogram', 'DATA'))
  582. copy_csv_files_to_root(rets_deepethogram_project_DATA_dir)
  583. retp_deepethogram_project_DATA_dir = os.path.abspath(os.path.join('.', 'ret', 'retp_deepethogram', 'DATA'))
  584. copy_csv_files_to_root(retp_deepethogram_project_DATA_dir)
  585. ret_labels_dir = os.path.abspath(os.path.join('.', 'ret', 'deg_csv_data'))
  586. temp_rets_labels_dir = os.path.abspath(os.path.join('.', 'ret', 'labels_csv_rets'))
  587. temp_retp_labels_dir = os.path.abspath(os.path.join('.', 'ret', 'labels_csv_retp'))
  588. move_csv_files_to_target_folders(rets_deepethogram_project_DATA_dir, temp_rets_labels_dir)
  589. move_csv_files_to_target_folders(retp_deepethogram_project_DATA_dir, temp_retp_labels_dir)
  590. rename_csv_files_remove_video_prefix(temp_rets_labels_dir)
  591. rename_csv_files_remove_video_prefix(temp_retp_labels_dir)
  592. # %%
  593. from convert_labelscsv2annot_ret import get_csv_files, read_labelscsv_file, create_annot_content, creat_annot_file
  594. ret_bento_dir = os.path.abspath(os.path.join('.', 'ret', 'bento_annot_data'))
  595. group = 's'
  596. csv_files = get_csv_files(temp_rets_labels_dir)
  597. for csv_file in csv_files:
  598. csv_file_path = os.path.join(temp_rets_labels_dir, csv_file)
  599. behavior_frames_dict = read_labelscsv_file(csv_file_path)
  600. annot_content = create_annot_content(csv_file)
  601. creat_annot_file(csv_file, ret_bento_dir, annot_content, behavior_frames_dict)
  602. group = 'p'
  603. csv_files = get_csv_files(temp_retp_labels_dir)
  604. for csv_file in csv_files:
  605. csv_file_path = os.path.join(temp_retp_labels_dir, csv_file)
  606. behavior_frames_dict = read_labelscsv_file(csv_file_path)
  607. annot_content = create_annot_content(csv_file)
  608. creat_annot_file(csv_file, ret_bento_dir, annot_content, behavior_frames_dict)
  609. # %%
  610. move_csv_files_to_target_folders(temp_rets_labels_dir, ret_labels_dir)
  611. move_csv_files_to_target_folders(temp_retp_labels_dir, ret_labels_dir)
  612. # %% [markdown]
  613. # ### Check ret overlap.
  614. # %%
  615. from collections import Counter
  616. # Control variable: if True, check all frames; if False, use specified frame range
  617. if_all_frame = False
  618. # Set frame range
  619. if if_all_frame:
  620. # Get the maximum frame number from all trials
  621. max_frame = 8900
  622. for group in ['s', 'p']:
  623. if group in ret_origin_bhvr_frames_dict:
  624. for rank in ['D', 'S']:
  625. if rank in ret_origin_bhvr_frames_dict[group]:
  626. for t_id in ret_origin_bhvr_frames_dict[group][rank].keys():
  627. trial_data = ret_origin_bhvr_frames_dict[group][rank][t_id]
  628. for behavior in trial_data.keys():
  629. if len(trial_data[behavior]) > 0:
  630. max_frame = max(max_frame, trial_data[behavior].index.max())
  631. frame_start = 0
  632. frame_end = max_frame + 1
  633. print("\n" + "="*80)
  634. print(f"Frame-by-Frame Behavior Combination Analysis (ALL frames: 0-{frame_end})")
  635. print("="*80)
  636. else:
  637. frame_start = 0
  638. frame_end = 8900
  639. print("\n" + "="*80)
  640. print("Frame-by-Frame Behavior Combination Analysis (0:8900 frames)")
  641. print("="*80)
  642. total_frames = frame_end - frame_start
  643. # Excluded behaviors
  644. excluded_behaviors = ['rat_in', 'in_proximity', 'stretching/rearing', 'stretching/rearing.1']
  645. # Iterate through all groups and ranks
  646. for group in ['s', 'p']:
  647. group_name = 'Single' if group == 's' else 'Pair'
  648. print(f"\n{'='*60}")
  649. print(f"Group: {group_name} (Frames: {frame_start}-{frame_end})")
  650. print('='*60)
  651. for rank in ['D', 'S']:
  652. rank_name = 'Dominant' if rank == 'D' else 'Subordinate'
  653. print(f"\n--- Rank: {rank_name} ---")
  654. if group not in ret_origin_bhvr_frames_dict or rank not in ret_origin_bhvr_frames_dict[group]:
  655. print(f" No data for {group_name} {rank_name}")
  656. continue
  657. # Initialize aggregated counter for all trials in this group-rank combination
  658. aggregated_combinations = []
  659. aggregated_overlap_frames = [] # Store all overlapping frames across trials
  660. trial_count = 0
  661. # Iterate through all trials
  662. for t_id in sorted(ret_origin_bhvr_frames_dict[group][rank].keys()):
  663. trial_data = ret_origin_bhvr_frames_dict[group][rank][t_id]
  664. # Get all available behaviors for this trial
  665. available_behaviors = [b for b in trial_data.keys() if b not in excluded_behaviors]
  666. if len(available_behaviors) == 0:
  667. print(f"\n {t_id}: No behaviors available after exclusion")
  668. continue
  669. # Create a list to store behavior combination for each frame
  670. frame_combinations = []
  671. overlap_frames_info = [] # Store frame numbers with overlapping behaviors
  672. # Iterate through each frame
  673. for frame_idx in range(frame_start, frame_end):
  674. # Find all active behaviors at this frame
  675. active_behaviors = []
  676. for behavior in sorted(available_behaviors): # Sort for consistent ordering
  677. try:
  678. if frame_idx in trial_data[behavior].index and trial_data[behavior][frame_idx] == 1:
  679. active_behaviors.append(behavior)
  680. except:
  681. pass
  682. # Create combination string
  683. if len(active_behaviors) == 0:
  684. combo_str = "no_behavior"
  685. elif len(active_behaviors) == 1:
  686. combo_str = active_behaviors[0]
  687. else:
  688. combo_str = "+".join(active_behaviors)
  689. # Record overlapping frame information
  690. overlap_frames_info.append({
  691. 'frame': frame_idx,
  692. 'behaviors': active_behaviors
  693. })
  694. frame_combinations.append(combo_str)
  695. # Add this trial's combinations to the aggregated list
  696. aggregated_combinations.extend(frame_combinations)
  697. # Add overlapping frames info with trial ID
  698. for frame_info in overlap_frames_info:
  699. aggregated_overlap_frames.append({
  700. 'trial': t_id,
  701. 'frame': frame_info['frame'],
  702. 'behaviors': frame_info['behaviors']
  703. })
  704. trial_count += 1
  705. # Count occurrences of each combination
  706. combination_counter = Counter(frame_combinations)
  707. # Calculate percentages and sort by frequency
  708. combination_stats = []
  709. for combo, count in combination_counter.items():
  710. pct = (count / total_frames) * 100
  711. combination_stats.append({
  712. 'combination': combo,
  713. 'frames': count,
  714. 'pct': pct
  715. })
  716. # Sort by frame count (descending)
  717. combination_stats.sort(key=lambda x: x['frames'], reverse=True)
  718. # Print results
  719. print(f"\n {t_id}:")
  720. print(f" Total unique behavior combinations: {len(combination_stats)}")
  721. print(f" Frame distribution:")
  722. # Show all combinations
  723. for stat in combination_stats:
  724. print(f" {stat['combination']}: {stat['frames']} frames ({stat['pct']:.2f}%)")
  725. # Display overlapping frames details
  726. if len(overlap_frames_info) > 0:
  727. print(f"\n Overlapping frames details ({len(overlap_frames_info)} frames total):")
  728. for frame_info in overlap_frames_info:
  729. behaviors_str = " + ".join(frame_info['behaviors'])
  730. print(f" Frame {frame_info['frame']}: {behaviors_str}")
  731. # Summary: single vs multiple behaviors
  732. single_behavior_frames = sum(stat['frames'] for stat in combination_stats
  733. if '+' not in stat['combination'] and stat['combination'] != 'no_behavior')
  734. multiple_behavior_frames = sum(stat['frames'] for stat in combination_stats
  735. if '+' in stat['combination'])
  736. no_behavior_frames = sum(stat['frames'] for stat in combination_stats
  737. if stat['combination'] == 'no_behavior')
  738. print(f"\n Summary:")
  739. print(f" Single behavior: {single_behavior_frames} frames ({single_behavior_frames/total_frames*100:.2f}%)")
  740. print(f" Multiple behaviors: {multiple_behavior_frames} frames ({multiple_behavior_frames/total_frames*100:.2f}%)")
  741. print(f" No behavior: {no_behavior_frames} frames ({no_behavior_frames/total_frames*100:.2f}%)")
  742. # Print aggregated statistics for this group-rank combination
  743. if trial_count > 0:
  744. print(f"\n{'*' * 50}")
  745. print(f"AGGREGATED STATISTICS: {group_name} - {rank_name}")
  746. print(f"Total trials: {trial_count}")
  747. print(f"Total frames analyzed: {len(aggregated_combinations)}")
  748. print(f"{'*' * 50}")
  749. # Count aggregated combinations
  750. aggregated_counter = Counter(aggregated_combinations)
  751. total_aggregated_frames = len(aggregated_combinations)
  752. # Calculate percentages and sort
  753. aggregated_stats = []
  754. for combo, count in aggregated_counter.items():
  755. pct = (count / total_aggregated_frames) * 100
  756. aggregated_stats.append({
  757. 'combination': combo,
  758. 'frames': count,
  759. 'pct': pct
  760. })
  761. aggregated_stats.sort(key=lambda x: x['frames'], reverse=True)
  762. print(f"\nAggregated behavior combinations:")
  763. for stat in aggregated_stats:
  764. print(f" {stat['combination']}: {stat['frames']} frames ({stat['pct']:.2f}%)")
  765. # Display all overlapping frames in aggregated data
  766. if len(aggregated_overlap_frames) > 0:
  767. print(f"\n All overlapping frames across trials ({len(aggregated_overlap_frames)} frames total):")
  768. # Group by trial for better readability
  769. current_trial = None
  770. for frame_info in aggregated_overlap_frames:
  771. if current_trial != frame_info['trial']:
  772. current_trial = frame_info['trial']
  773. print(f"\n {current_trial}:")
  774. behaviors_str = " + ".join(frame_info['behaviors'])
  775. print(f" Frame {frame_info['frame']}: {behaviors_str}")
  776. # Aggregated summary
  777. agg_single = sum(stat['frames'] for stat in aggregated_stats
  778. if '+' not in stat['combination'] and stat['combination'] != 'no_behavior')
  779. agg_multiple = sum(stat['frames'] for stat in aggregated_stats
  780. if '+' in stat['combination'])
  781. agg_no_behavior = sum(stat['frames'] for stat in aggregated_stats
  782. if stat['combination'] == 'no_behavior')
  783. print(f"\nAggregated Summary:")
  784. print(f" Single behavior: {agg_single} frames ({agg_single/total_aggregated_frames*100:.2f}%)")
  785. print(f" Multiple behaviors: {agg_multiple} frames ({agg_multiple/total_aggregated_frames*100:.2f}%)")
  786. print(f" No behavior: {agg_no_behavior} frames ({agg_no_behavior/total_aggregated_frames*100:.2f}%)")
  787. print(f"{'*' * 50}\n")
  788. print("\n" + "="*80)
  789. print("Frame-by-frame analysis complete")
  790. print("="*80)
  791. # %% [markdown]
  792. # ## Analyze data and plot figures.
  793. # %% [markdown]
  794. # ### lst: lot ethogram and velocity
  795. # %%
  796. # Plot ethogram and velocity for all trials
  797. import matplotlib.pyplot as plt
  798. import matplotlib.patches as mpatches
  799. import os
  800. # Create output directory if it doesn't exist
  801. os.makedirs('data/ethogram_velocity_plots', exist_ok=True)
  802. # Behavior colors
  803. behavior_colors = {'escape': 'red', 'freezing': 'blue', 'sniffing': 'green',
  804. 'grooming': 'purple', 'stretching': 'orange', 'rearing': 'brown'}
  805. # Loop through all trials
  806. plot_count = 0
  807. for g in ['s', 'p']:
  808. for r in ['D', 'S']:
  809. for t_id in lst_bhvr_frames_dict[g][r].keys():
  810. try:
  811. # Parse m_id from t_id
  812. if t_id.split('B')[1].split('M')[0] in ['1', '3', '7', '10', '11', '12', '13', '14', '15', '16', '17', '20']:
  813. m = 1 if t_id.split('M')[1].split('T')[0] == 'S' else 2
  814. else:
  815. m = 1 if t_id.split('M')[1].split('T')[0] == 'D' else 2
  816. m_id = t_id.split('M')[0] + 'M' + str(m) + '(' + t_id.split('M')[1].split('T')[0] + ')'
  817. # Get velocity data and looming_start_frame based on group
  818. if g == 's':
  819. velocity_data = lsts_data[m_id]['velocity']['waist']
  820. looming_start_frame = looming_start_frames[m_id][int(t_id.split('T')[1])-1]
  821. elif g == 'p':
  822. velocity_data = lstp_data[m_id]['velocity']['waist']
  823. mp_id = mp_ids[int(m_id.split('B')[1].split('M')[0])-1]
  824. looming_start_frame = looming_start_frames[mp_id][int(t_id.split('T')[1])-1]
  825. # Get behavior data for this trial
  826. trial_bhvr_data = lst_bhvr_frames_dict[g][r][t_id]
  827. # Define time window: -5s to +60s (150 frames before to 1800 frames after looming start)
  828. plot_start_frame = looming_start_frame - 150
  829. plot_end_frame = looming_start_frame + 1800
  830. plot_window = slice(plot_start_frame, plot_end_frame)
  831. # Extract velocity data for plotting (convert to cm/s)
  832. velocity_segment = velocity_data[plot_window].values
  833. velocity_cm_s = velocity_segment / lst_pixpercm * framerate
  834. time_axis = np.arange(len(velocity_cm_s)) / framerate # Convert to seconds
  835. # Get all behaviors
  836. behaviors = [b for b in trial_bhvr_data.keys() if b != 'nest']
  837. # Create figure with two subplots
  838. fig, (ax1, ax2) = plt.subplots(2, 1, figsize=(16, 8), sharex=True,
  839. gridspec_kw={'height_ratios': [1, 2]})
  840. # Plot 1: Velocity
  841. ax1.plot(time_axis, velocity_cm_s, 'k-', linewidth=1.5, label='Velocity')
  842. ax1.axvline(x=150/framerate, color='red', linestyle='--', linewidth=2, label='Looming start')
  843. ax1.set_ylabel('Velocity (cm/s)', fontsize=12, fontweight='bold')
  844. ax1.set_title(f'Trial: {t_id} ({m_id}) - Group: {g}, Rank: {r}', fontsize=14, fontweight='bold')
  845. ax1.legend(loc='upper right')
  846. ax1.grid(True, alpha=0.3)
  847. ax1.set_ylim(bottom=0)
  848. # Plot 2: Ethogram
  849. y_pos = 0
  850. behavior_positions = {}
  851. for behavior in behaviors:
  852. behavior_positions[behavior] = y_pos
  853. color = behavior_colors.get(behavior, 'gray')
  854. # Get behavior frames (relative to plot window)
  855. for frame_idx in range(plot_start_frame, plot_end_frame):
  856. relative_frame = frame_idx - looming_start_frame + 150 # Relative to looming window
  857. try:
  858. if relative_frame in trial_bhvr_data[behavior].index:
  859. if trial_bhvr_data[behavior][relative_frame] == 1:
  860. time_point = (frame_idx - plot_start_frame) / framerate
  861. ax2.barh(y_pos, 1/framerate, left=time_point, height=0.8,
  862. color=color, edgecolor='none', alpha=0.8)
  863. except:
  864. pass
  865. y_pos += 1
  866. # Add vertical line for looming start
  867. ax2.axvline(x=150/framerate, color='red', linestyle='--', linewidth=2, label='Looming start')
  868. ax2.set_yticks(range(len(behaviors)))
  869. ax2.set_yticklabels(behaviors, fontsize=10)
  870. ax2.set_xlabel('Time (s)', fontsize=12, fontweight='bold')
  871. ax2.set_ylabel('Behaviors', fontsize=12, fontweight='bold')
  872. ax2.set_xlim(0, (plot_end_frame - plot_start_frame) / framerate)
  873. ax2.grid(True, axis='x', alpha=0.3)
  874. # Add legend for behaviors
  875. legend_patches = [mpatches.Patch(color=behavior_colors.get(b, 'gray'), label=b)
  876. for b in behaviors if b in behavior_colors]
  877. ax2.legend(handles=legend_patches, loc='upper right', ncol=3, fontsize=9)
  878. plt.tight_layout()
  879. plt.savefig(f'data/ethogram_velocity_plots/lst{g}_{t_id}_ethogram_velocity.png', dpi=300, bbox_inches='tight')
  880. plt.close()
  881. plot_count += 1
  882. if plot_count % 10 == 0:
  883. print(f"Processed {plot_count} trials...")
  884. except Exception as e:
  885. print(f"Error processing trial {t_id}: {str(e)}")
  886. continue
  887. print(f"\nCompleted! Generated {plot_count} plots.")
  888. print(f"All plots saved in: data/ethogram_velocity_plots/")
  889. print(f"Time window for each plot: -5s to +60s relative to looming start")
  890. # %% [markdown]
  891. # ### ret: plot ethogram
  892. # %%
  893. # Plot ethogram for ret trials with x-coordinate trajectory
  894. import matplotlib.pyplot as plt
  895. import matplotlib.patches as mpatches
  896. import os
  897. import pandas as pd
  898. import numpy as np
  899. from analyze_data_utils import ret_extract_target_data
  900. # Create output directory if it doesn't exist
  901. os.makedirs('data/ret_ethogram_plots', exist_ok=True)
  902. # Extract x-coordinate data for all trials
  903. x_coordinate_data = ret_extract_target_data(ret_raw_data, target_key=['X center'], times=['withrat'])
  904. # Behavior colors for ret
  905. behavior_colors_ret = {
  906. 'rat_in': 'black',
  907. 'approach': 'red',
  908. 'investigation': 'orange',
  909. 'withdrawal': 'blue',
  910. 'stretch-attend': 'cyan',
  911. 'freezing': 'navy',
  912. 'tail_rattling': 'purple',
  913. 'huddling': 'pink',
  914. 'approach_partner': 'lightcoral',
  915. 'follow_partner': 'salmon',
  916. 'groom_partner': 'gold',
  917. 'sniff_partner': 'yellow',
  918. 'grooming': 'green'
  919. }
  920. # Loop through all trials
  921. plot_count = 0
  922. for g in ['s', 'p']:
  923. for r in ['D', 'S']:
  924. if g not in ret_bhvr_frames_dict or r not in ret_bhvr_frames_dict[g]:
  925. continue
  926. for t_id in ret_bhvr_frames_dict[g][r].keys():
  927. # Get behavior data for this trial
  928. trial_bhvr_data = ret_bhvr_frames_dict[g][r][t_id]
  929. if g == 's':
  930. if t_id[1] in ['4']:
  931. if t_id[-1] == 'S':
  932. m_id = '1'
  933. elif t_id[-1] == 'D':
  934. m_id = '2'
  935. else:
  936. if t_id[-1] == 'D':
  937. m_id = '1'
  938. elif t_id[-1] == 'S':
  939. m_id = '2'
  940. sess_id = t_id.split('M')[0]+'M'+m_id
  941. elif g =='p':
  942. sess_id = t_id.split('M')[0]+'M1&M2'
  943. # Extract x-coordinate data for this session
  944. session_label = t_id.split('M')[0] # e.g., 'D1'
  945. cond = 'MD' if r == 'D' else 'MS'
  946. # Get x-coordinate series
  947. if g == 's':
  948. x_series = x_coordinate_data['single'][cond]['withrat'][session_label]
  949. else:
  950. x_series = x_coordinate_data['pair'][cond]['withrat'][session_label]
  951. # Convert to numeric and handle any non-numeric values
  952. x_values = pd.to_numeric(x_series, errors='coerce')
  953. # Calculate the offset between behavior and x-coordinate data
  954. offset_frames = rat_in_frames[sess_id] - offset_frames[sess_id]
  955. # Behavior data range (always 0 to 8900)
  956. behavior_start = 0
  957. behavior_end = 8900
  958. # Extract x-coordinate data with offset handling
  959. if offset_frames >= 0:
  960. # Normal case: x-coordinate starts before or at behavior start
  961. plot_x_values = x_values[offset_frames:offset_frames + 8900]
  962. x_time_offset = 0
  963. else:
  964. # offset_frames is negative: behavior starts before x-coordinate data
  965. # Need to pad the beginning with NaN and shift x data to the right
  966. x_start = 0
  967. x_end = 8900 + offset_frames # This will be less than 8900
  968. plot_x_values = pd.Series([np.nan] * (-offset_frames)).append(
  969. x_values[x_start:x_end], ignore_index=True)
  970. x_time_offset = -offset_frames
  971. # Get all behaviors
  972. behaviors = [b for b in trial_bhvr_data.keys() if b != 'rat_in']
  973. # Create figure with two subplots (x-coordinate and ethogram)
  974. fig, (ax1, ax2) = plt.subplots(2, 1, figsize=(16, 10), sharex=True,
  975. gridspec_kw={'height_ratios': [1, 2]})
  976. # Plot 1: X-coordinate trajectory
  977. time_axis = np.arange(len(plot_x_values)) / framerate
  978. ax1.plot(time_axis, plot_x_values, color='black', linewidth=0.5, alpha=0.7)
  979. ax1.set_ylabel('X Coordinate (cm)', fontsize=12, fontweight='bold')
  980. ax1.set_title(f'RET Trial: {t_id} - Group: {g}, Rank: {r}', fontsize=14, fontweight='bold')
  981. ax1.grid(True, axis='both', alpha=0.3)
  982. # Add vertical lines for time markers on x-coordinate plot
  983. ax1.axvline(x=0/framerate, color='red', linestyle='--', linewidth=2, alpha=0.5)
  984. ax1.axvline(x=1800/framerate, color='red', linestyle='--', linewidth=2, alpha=0.5)
  985. ax1.axvline(x=3600/framerate, color='red', linestyle='--', linewidth=2, alpha=0.5)
  986. ax1.axvline(x=5400/framerate, color='red', linestyle='--', linewidth=2, alpha=0.5)
  987. ax1.axvline(x=7200/framerate, color='red', linestyle='--', linewidth=2, alpha=0.5)
  988. ax1.axvline(x=8900/framerate, color='red', linestyle='--', linewidth=2, alpha=0.5)
  989. # Plot 2: Ethogram
  990. y_pos = 0
  991. behavior_positions = {}
  992. for behavior in behaviors:
  993. behavior_positions[behavior] = y_pos
  994. color = behavior_colors_ret.get(behavior, 'gray')
  995. # Get behavior frames (behavior data is always from 0 to 8900)
  996. for frame_idx in range(behavior_start, behavior_end):
  997. try:
  998. if frame_idx in trial_bhvr_data[behavior].index:
  999. if trial_bhvr_data[behavior][frame_idx] == 1:
  1000. time_point = frame_idx / framerate
  1001. ax2.barh(y_pos, 1/framerate, left=time_point, height=0.8,
  1002. color=color, edgecolor='none', alpha=0.8)
  1003. except:
  1004. pass
  1005. y_pos += 1
  1006. # Add vertical lines for time markers on ethogram
  1007. ax2.axvline(x=0/framerate, color='red', linestyle='--', linewidth=2)
  1008. ax2.axvline(x=1800/framerate, color='red', linestyle='--', linewidth=2)
  1009. ax2.axvline(x=3600/framerate, color='red', linestyle='--', linewidth=2)
  1010. ax2.axvline(x=5400/framerate, color='red', linestyle='--', linewidth=2)
  1011. ax2.axvline(x=7200/framerate, color='red', linestyle='--', linewidth=2)
  1012. ax2.axvline(x=8900/framerate, color='red', linestyle='--', linewidth=2)
  1013. ax2.set_yticks(range(len(behaviors)))
  1014. ax2.set_yticklabels(behaviors, fontsize=10)
  1015. ax2.set_xlabel('Time (s)', fontsize=12, fontweight='bold')
  1016. ax2.set_ylabel('Behaviors', fontsize=12, fontweight='bold')
  1017. ax2.set_xlim(0, 8900 / framerate)
  1018. ax2.grid(True, axis='x', alpha=0.3)
  1019. # Add legend for behaviors
  1020. legend_patches = [mpatches.Patch(color=behavior_colors_ret.get(b, 'gray'), label=b)
  1021. for b in behaviors if b in behavior_colors_ret]
  1022. # ax2.legend(handles=legend_patches, loc='upper right', ncol=4, fontsize=8)
  1023. plt.tight_layout()
  1024. plt.savefig(f'data/ret_ethogram_plots/ret{g}_{t_id}_ethogram_with_x.png', dpi=300, bbox_inches='tight')
  1025. plt.close()
  1026. plot_count += 1
  1027. if plot_count % 10 == 0:
  1028. print(f"Processed {plot_count} trials...")
  1029. print(f"\nCompleted! Generated {plot_count} plots.")
  1030. print(f"All plots saved in: data/ret_ethogram_plots/")
  1031. print(f"Time window for each plot: 0 - 5 min")
  1032. # %% [markdown]
  1033. # ## Fig1
  1034. # %% [markdown]
  1035. # ### Fig1B&D. location and velocity histogram
  1036. # %%
  1037. import numpy as np
  1038. import matplotlib.pyplot as plt
  1039. from matplotlib.colors import ListedColormap, BoundaryNorm
  1040. import seaborn as sns
  1041. from analyze_data_utils import filter_in_range, filter_dict_data
  1042. import matplotlib.colors as mcolors
  1043. import pandas as pd
  1044. nbins = 25
  1045. cmap = 'viridis'
  1046. pixpergrid = 43.52
  1047. pixpercm = pixpergrid / 2 # pixels per cm
  1048. def prepare_velocity_data(ids, data_source, min_speed=0, max_speed=300):
  1049. all_x = []
  1050. all_y = []
  1051. all_vx = []
  1052. all_vy = []
  1053. all_ids = []
  1054. for mouse_id in ids:
  1055. # interested_frame = (looming_start_frames[mouse_id][0], looming_start_frames[mouse_id][-1]+1800)
  1056. # interested_frame = [0, 18000]
  1057. # x = data_source[mouse_id]['coordinate']['waist']['x'][interested_frame[0]:interested_frame[1]]
  1058. # y = data_source[mouse_id]['coordinate']['waist']['y'][interested_frame[0]:interested_frame[1]]
  1059. # # Calculate velocity vectors (cm/s)
  1060. # vx = np.diff(x) / pixpercm * framerate
  1061. # vy = np.diff(y) / pixpercm * framerate
  1062. # # Keep only points with speed >= min_speed
  1063. # speed = np.sqrt(vx**2 + vy**2)
  1064. # mask = (speed >= min_speed) & (speed <= max_speed)
  1065. # all_x.extend(x[:-1][mask]) # x[:-1] aligns with diff length
  1066. # all_y.extend(y[:-1][mask])
  1067. # all_vx.extend(vx[mask])
  1068. # all_vy.extend(vy[mask])
  1069. # all_ids.extend([mouse_id] * np.sum(mask)) # Record mouse ID for each point
  1070. start_frames = looming_start_frames[mouse_id]
  1071. for start_frame in start_frames:
  1072. start_frame = start_frame + 30 * 20
  1073. end_frame = start_frame + 30 * 60
  1074. max_frame = len(data_source[mouse_id]['coordinate']['waist']['x'])
  1075. end_frame = min(end_frame, max_frame)
  1076. x = data_source[mouse_id]['coordinate']['waist']['x'][start_frame:end_frame]
  1077. y = data_source[mouse_id]['coordinate']['waist']['y'][start_frame:end_frame]
  1078. vx = np.diff(x) / pixpercm * framerate
  1079. vy = np.diff(y) / pixpercm * framerate
  1080. speed = np.sqrt(vx**2 + vy**2)
  1081. mask = (speed >= min_speed) & (speed <= max_speed)
  1082. all_x.extend(x[:-1][mask])
  1083. all_y.extend(y[:-1][mask])
  1084. all_vx.extend(vx[mask])
  1085. all_vy.extend(vy[mask])
  1086. all_ids.extend([mouse_id] * np.sum(mask))
  1087. return np.array(all_x), np.array(all_y), np.array(all_vx), np.array(all_vy), np.array(all_ids)
  1088. # def prepare_velocity_data(ids, data_source, min_speed=0, max_speed=300):
  1089. # all_x = []
  1090. # all_y = []
  1091. # all_vx = []
  1092. # all_vy = []
  1093. # all_ids = []
  1094. # for mouse_id in ids:
  1095. # # Get the entire time period [first stimulus start, last stimulus end+1800]
  1096. # start_frame0 = looming_start_frames[mouse_id][0]
  1097. # end_frame_last = looming_start_frames[mouse_id][-1] + 1800
  1098. # max_frame = len(data_source[mouse_id]['coordinate']['waist']['x'])
  1099. # end_frame_last = min(end_frame_last, max_frame) # Ensure not exceeding data range
  1100. # # Extract coordinate data for the entire time period
  1101. # x_full = data_source[mouse_id]['coordinate']['waist']['x'][start_frame0:end_frame_last]
  1102. # y_full = data_source[mouse_id]['coordinate']['waist']['y'][start_frame0:end_frame_last]
  1103. # # Create a mask initialized to True, indicating all points are initially included
  1104. # mask = np.ones(len(x_full), dtype=bool)
  1105. # # Iterate through each stimulus event, excluding data during events
  1106. # for start_frame in looming_start_frames[mouse_id]:
  1107. # # Calculate the relative position of the current event in the full data
  1108. # event_start = start_frame - start_frame0
  1109. # event_end = min(event_start + 1800, len(x_full))
  1110. # # Mark data during the event as False (excluded)
  1111. # if event_start < len(mask):
  1112. # mask[event_start:event_end] = False
  1113. # # Apply mask to exclude all data during stimulus events
  1114. # x = x_full[mask]
  1115. # y = y_full[mask]
  1116. # # Calculate velocity vectors (cm/s)
  1117. # vx = np.diff(x) / pixpercm * framerate
  1118. # vy = np.diff(y) / pixpercm * framerate
  1119. # # Keep only points with speed in [min_speed, max_speed]
  1120. # speed = np.sqrt(vx**2 + vy**2)
  1121. # speed_mask = (speed >= min_speed) & (speed <= max_speed)
  1122. # # Add data to total lists
  1123. # all_x.extend(x[:-1][speed_mask]) # x[:-1] aligns with diff length
  1124. # all_y.extend(y[:-1][speed_mask])
  1125. # all_vx.extend(vx[speed_mask])
  1126. # all_vy.extend(vy[speed_mask])
  1127. # all_ids.extend([mouse_id] * np.sum(speed_mask))
  1128. # return np.array(all_x), np.array(all_y), np.array(all_vx), np.array(all_vy), np.array(all_ids)
  1129. def prepare_nest_velocity(ids, data_source, nest_data, min_speed=5, max_speed=300):
  1130. all_x = []
  1131. all_y = []
  1132. all_vx = []
  1133. all_vy = []
  1134. all_ids = []
  1135. for mouse_id in ids:
  1136. interested_frame = (looming_start_frames[mouse_id][0], looming_start_frames[mouse_id][-1]+1800)
  1137. # interested_frame = [0, 18000]
  1138. nest_id = mouse_id.split('M')[0]+'M'+mouse_id[-2]+'T1'
  1139. g = 's'
  1140. r = mouse_id[-2]
  1141. start_frames = looming_start_frames[mouse_id]
  1142. for start_frame in start_frames:
  1143. # start_frame = 9000
  1144. if nest_id in nest_data[g][r].keys():
  1145. nest_list = nest_data[g][r][nest_id]
  1146. else:
  1147. continue
  1148. for nest in nest_list:
  1149. if np.any(np.isnan(nest)):
  1150. continue
  1151. else:
  1152. start, end = nest
  1153. nest_start_frame = start_frame + start
  1154. # interested_frame = [nest_start_frame-60, nest_start_frame+10] # baseline nest return
  1155. interested_frame = [nest_start_frame-150-60, nest_start_frame-150+10] # looming nest return
  1156. # interested_frame = [start_frame+start-150-10, start_frame+end-150+10] # escape
  1157. x = data_source[mouse_id]['coordinate']['waist']['x'][interested_frame[0]:interested_frame[1]]
  1158. y = data_source[mouse_id]['coordinate']['waist']['y'][interested_frame[0]:interested_frame[1]]
  1159. # Calculate velocity vectors (cm/s)
  1160. vx = np.diff(x) / pixpercm * framerate
  1161. vy = np.diff(y) / pixpercm * framerate
  1162. # Keep only points with speed >= min_speed
  1163. speed = np.sqrt(vx**2 + vy**2)
  1164. mask = (speed >= min_speed) & (speed <= max_speed)
  1165. all_x.extend(x[:-1][mask]) # x[:-1] aligns with diff length
  1166. all_y.extend(y[:-1][mask])
  1167. all_vx.extend(vx[mask])
  1168. all_vy.extend(vy[mask])
  1169. all_ids.extend([mouse_id] * np.sum(mask)) # Record mouse ID for each point
  1170. return np.array(all_x), np.array(all_y), np.array(all_vx), np.array(all_vy), np.array(all_ids)
  1171. # all_speed_avg = []
  1172. # for m_id in ms_ids:
  1173. # nest_data = filter_dict_data(baseline_behavior_frames_dict, 'nest')
  1174. nest_data = filter_dict_data(lst_baseline_bhvr_labels_dict, 'nest')
  1175. nest_data = filter_in_range(nest_data, t_range=(0,18000), method='replace')
  1176. # x, y, vx, vy, point_ids = prepare_nest_velocity(ms_ids, lsts_data, nest_data)
  1177. x, y, vx, vy, point_ids = prepare_velocity_data(ms_ids, lsts_data)
  1178. heatmap, xedges, yedges = np.histogram2d(x, y, range=[[0, 1088], [0, 1088]], bins=nbins)
  1179. mouse_count_grid = np.zeros_like(heatmap, dtype=int)
  1180. x_bin_indices = np.digitize(x, xedges) - 1
  1181. y_bin_indices = np.digitize(y, yedges) - 1
  1182. cell_mice_dict = {}
  1183. for xi, yi, mouse_id in zip(x_bin_indices, y_bin_indices, point_ids):
  1184. if 0 <= xi < nbins and 0 <= yi < nbins:
  1185. key = (xi, yi)
  1186. if key not in cell_mice_dict:
  1187. cell_mice_dict[key] = set()
  1188. cell_mice_dict[key].add(mouse_id)
  1189. for (xi, yi), mice_set in cell_mice_dict.items():
  1190. mouse_count_grid[xi, yi] = len(mice_set)
  1191. # mask = mouse_count_grid >= 3
  1192. grid_vx, _, _ = np.histogram2d(x, y, range=[[0, 1088], [0, 1088]], bins=nbins, weights=vx)
  1193. grid_vy, _, _ = np.histogram2d(x, y, range=[[0, 1088], [0, 1088]], bins=nbins, weights=vy)
  1194. count, _, _ = np.histogram2d(x, y, range=[[0, 1088], [0, 1088]], bins=nbins)
  1195. # count[count == 0] = 1
  1196. grid_vx /= count
  1197. grid_vy /= count
  1198. grid_vx /= 2
  1199. grid_vy /= 2 # change scale for nest return and escape
  1200. mask = (mouse_count_grid >= 3) & (count >= 5)
  1201. grid_vx[~mask] = 1e-5
  1202. grid_vy[~mask] = 1e-5
  1203. speed_grid = np.sqrt(grid_vx**2 + grid_vy**2)
  1204. speeds = np.sqrt(vx**2 + vy**2)
  1205. speed_sum, _, _ = np.histogram2d(x, y, range=[[0, 1088], [0, 1088]], bins=[xedges, yedges], weights=speeds)
  1206. speed_avg = speed_sum / count
  1207. speed_avg[~mask] = 1e-5
  1208. # all_speed_avg.append(speed_avg.T)
  1209. X, Y = np.meshgrid((xedges[:-1] + xedges[1:]) / 2,
  1210. (yedges[:-1] + yedges[1:]) / 2)
  1211. hist_frequency = heatmap / np.sum(heatmap) * 100
  1212. hist_frequency[~mask] = 1e-5
  1213. log_heatmap = np.log1p(hist_frequency)
  1214. magnitude = np.sqrt(grid_vx**2 + grid_vy**2)
  1215. scaled_magnitude = np.sqrt(magnitude + 1e-5)
  1216. vx_mod = grid_vx / (magnitude + 1e-5) * scaled_magnitude
  1217. vy_mod = grid_vy / (magnitude + 1e-5) * scaled_magnitude
  1218. # Save data for plotting
  1219. df = pd.DataFrame({
  1220. 'x_index': np.repeat(np.arange(nbins), nbins),
  1221. 'y_index': np.tile(np.arange(nbins), nbins),
  1222. 'time_pct': hist_frequency.T.flatten(),
  1223. 'speed_avg': speed_avg.T.flatten(),
  1224. 'vx': grid_vx.T.flatten(),
  1225. 'vy': grid_vy.T.flatten()
  1226. })
  1227. df.to_csv('data/Fig1B&D_location_speed_heatmap_60s.csv', index=False)
  1228. # %%
  1229. import numpy as np
  1230. import matplotlib.pyplot as plt
  1231. from matplotlib.colors import ListedColormap, BoundaryNorm
  1232. import seaborn as sns
  1233. import matplotlib.colors as mcolors
  1234. import pandas as pd
  1235. nbins = 25
  1236. df = pd.read_csv('data/Fig1B&D_location_speed_heatmap_60s.csv')
  1237. # Reshape to an (nbins, nbins) matrix (no transpose, matching the shape of the original computed variable).
  1238. hist_frequency = df.pivot(index='y_index', columns='x_index', values='time_pct').values
  1239. speed_avg = df.pivot(index='y_index', columns='x_index', values='speed_avg').values
  1240. grid_vx = df.pivot(index='y_index', columns='x_index', values='vx').values
  1241. grid_vy = df.pivot(index='y_index', columns='x_index', values='vy').values
  1242. xedges = np.linspace(0, nbins, nbins+1)
  1243. yedges = xedges.copy()
  1244. fig, ax = plt.subplots(figsize=(5,5), dpi=200)
  1245. # 1) draw coordinate heatmap
  1246. im = ax.imshow(
  1247. hist_frequency.T,
  1248. cmap=cmap,
  1249. origin='lower', # Start drawing from bottom-left corner
  1250. extent=(0, nbins, 0, nbins), # Map x and y axes to [0, nbins]
  1251. vmin=0,
  1252. vmax=4
  1253. )
  1254. fig.colorbar(im, label='Time (%)')
  1255. line_cor = 0.3
  1256. ax.plot([10*nbins/25+line_cor, 10*nbins/25+line_cor], [17.5*nbins/25-line_cor, 25*nbins/25-line_cor], color='white', linewidth=2)
  1257. ax.plot([0+line_cor, 0+line_cor], [17.5*nbins/25-line_cor, 25*nbins/25-line_cor], color='white', linewidth=2)
  1258. ax.plot([0+line_cor, 10*nbins/25+line_cor], [25*nbins/25-line_cor, 25*nbins/25-line_cor], color='white', linewidth=2)
  1259. ax.set_xlim(0, nbins)
  1260. ax.set_ylim(0, nbins)
  1261. ax.set_xticks([])
  1262. ax.set_yticks([])
  1263. ax.invert_yaxis()
  1264. plt.tight_layout()
  1265. # plt.savefig(f'G:/Li_lab/ppt/S_paper/Paper_v7/new_figures/coord_{title}.eps', format="eps", dpi=300, bbox_inches="tight")
  1266. plt.show()
  1267. fig, ax = plt.subplots(figsize=(5,5), dpi=300)
  1268. # 1) draw speed heatmap
  1269. plt.imshow(speed_avg.T, cmap=cmap, origin='lower',
  1270. extent=(0, nbins, 0, nbins), vmin=0, vmax=40)
  1271. plt.colorbar(label='Speed (cm/s)')
  1272. # ax.plot([10*nbins/25+line_cor, 10*nbins/25+line_cor], [17.5*nbins/25-line_cor, 25*nbins/25-line_cor], color='white', linewidth=2)
  1273. # ax.plot([0+line_cor, 0+line_cor], [17.5*nbins/25-line_cor, 25*nbins/25-line_cor], color='white', linewidth=2)
  1274. # ax.plot([0+line_cor, 10*nbins/25+line_cor], [25*nbins/25-line_cor, 25*nbins/25-line_cor], color='white', linewidth=2)
  1275. # ax.set_xticks([])
  1276. # ax.set_yticks([])
  1277. # ax.invert_yaxis()
  1278. # plt.savefig(f'G:/Li_lab/ppt/S_paper/Paper_v5/Figures_mod_v2/speed_heatmap_{title}.eps', format="eps", dpi=300, bbox_inches="tight")
  1279. # plt.show()
  1280. # fig, ax = plt.subplots(figsize=(5,5), dpi=200)
  1281. # ax.set_facecolor('black') # Plot area background
  1282. X, Y = np.meshgrid(np.arange(nbins) + 0.5, np.arange(nbins) + 0.5)
  1283. scale = 0.5
  1284. head_scale = 0.1
  1285. ax.set_xlim(X.min(), X.max())
  1286. ax.set_ylim(Y.min(), Y.max())
  1287. for i in range(0, X.shape[0], 1):
  1288. for j in range(0, X.shape[1], 1):
  1289. vx_ = grid_vx.T[i, j]
  1290. vy_ = grid_vy.T[i, j]
  1291. mag = np.sqrt(vx_**2 + vy_**2)
  1292. dx = vx_ * scale
  1293. dy = vy_ * scale
  1294. # x0 = X[i, j] - dx / 2
  1295. # y0 = Y[i, j] - dy / 2
  1296. # x1 = X[i, j] + dx / 2
  1297. # y1 = Y[i, j] + dy / 2
  1298. x0 = X[i, j]
  1299. y0 = Y[i, j]
  1300. x1 = x0 + dx
  1301. y1 = y0 + dy
  1302. ax.annotate('', xy=(x1, y1), xytext=(x0, y0),
  1303. arrowprops=dict(
  1304. arrowstyle='->, head_width={:.2f}, head_length={:.2f}'.format(mag * head_scale, mag * head_scale * 1.5),
  1305. color='white',
  1306. linewidth=1,
  1307. mutation_scale=5,
  1308. shrinkA=0, shrinkB=0
  1309. ))
  1310. ax.invert_yaxis()
  1311. line_cor = 0.3
  1312. ax.plot([10*nbins/25+line_cor, 10*nbins/25+line_cor], [17.5*nbins/25-line_cor, 25*nbins/25-line_cor], color='white', linewidth=2)
  1313. ax.plot([0+line_cor, 0+line_cor], [17.5*nbins/25-line_cor, 25*nbins/25-line_cor], color='white', linewidth=2)
  1314. ax.plot([0+line_cor, 10*nbins/25+line_cor], [25*nbins/25-line_cor, 25*nbins/25-line_cor], color='white', linewidth=2)
  1315. ax.set_xlim(0, nbins)
  1316. ax.set_ylim(0, nbins)
  1317. ax.set_xticks([])
  1318. ax.set_yticks([])
  1319. ax.invert_yaxis()
  1320. plt.tight_layout()
  1321. # plt.savefig(f'G:/Li_lab/ppt/S_paper/Paper_v7/new_figures/speed_{title}.eps', format="eps", dpi=300, bbox_inches="tight")
  1322. plt.show()
  1323. # %% [markdown]
  1324. # ### Fig1C&E & S1B: Time in zone, Speed in zone.
  1325. # %%
  1326. import numpy as np
  1327. import pandas as pd
  1328. from analyze_data_utils import get_lst_location_value
  1329. # Define parameters
  1330. framerate = 30 # fps
  1331. pixpergrid = 43.52
  1332. pixpercm = pixpergrid / 2 # pixels per cm
  1333. # Define time periods (in frames)
  1334. time_periods = {
  1335. 'baseline': (0, 10*60*framerate), # baseline 10 minutes
  1336. '0-5s': (0, 5 * framerate), # 0-5 seconds after looming
  1337. '5-20s': (5 * framerate, 20 * framerate), # 5-20 seconds after looming
  1338. '20-60s': (20 * framerate, 60 * framerate) # 20-60 seconds after looming
  1339. }
  1340. # Define zone labels
  1341. zone_labels = {0: 'nest', 1: 'edge', 2: 'center'}
  1342. # zones_time_pcts[group][rank][time_period][zone] = [percentages for each mouse]
  1343. zones_time_pcts = {'s': {'D': {}, 'S': {}}, 'p': {'D': {}, 'S': {}}}
  1344. zones_speed = {'s': {'D': {}, 'S': {}}, 'p': {'D': {}, 'S': {}}}
  1345. for g in ['s', 'p']:
  1346. for m in ['D', 'S']:
  1347. for period in time_periods.keys():
  1348. zones_time_pcts[g][m][period] = {0: [], 1: [], 2: []} # nest, edge, center
  1349. zones_speed[g][m][period] = {0: [], 1: [], 2: []}
  1350. for g in ['s', 'p']:
  1351. group_name = 'single' if g == 's' else 'pair'
  1352. for m in ['D', 'S']:
  1353. rank_name = 'Dominant' if m == 'D' else 'Subordinate'
  1354. m_ids = mS_ids if m == 'S' else mD_ids
  1355. lst_data = lsts_data if g == 's' else lstp_data
  1356. for m_id in m_ids:
  1357. for period_name, (start_offset, end_offset) in time_periods.items():
  1358. if period_name == 'baseline':
  1359. interested_frame = (start_offset, end_offset)
  1360. coord_data = {
  1361. 'x': lst_data[m_id]['coordinate']['waist']['x'][interested_frame[0]:interested_frame[1]].reset_index(drop=True),
  1362. 'y': lst_data[m_id]['coordinate']['waist']['y'][interested_frame[0]:interested_frame[1]].reset_index(drop=True)
  1363. }
  1364. location_value = get_lst_location_value(coord_data)
  1365. velocity_value = lst_data[m_id]['velocity']['waist'][interested_frame[0]:interested_frame[1]].reset_index(drop=True)
  1366. for zone in [0, 1, 2]:
  1367. zone_count = (location_value == zone).sum()
  1368. zone_pct = zone_count / len(location_value) * 100
  1369. zones_time_pcts[g][m][period_name][zone].append(zone_pct)
  1370. zone_speed = velocity_value[location_value == zone].mean() / pixpercm * framerate
  1371. zones_speed[g][m][period_name][zone].append(zone_speed)
  1372. else:
  1373. for lsf in looming_start_frames.get(m_id, []):
  1374. interested_frame = (lsf + start_offset, lsf + end_offset)
  1375. coord_data = {
  1376. 'x': lst_data[m_id]['coordinate']['waist']['x'][interested_frame[0]:interested_frame[1]].reset_index(drop=True),
  1377. 'y': lst_data[m_id]['coordinate']['waist']['y'][interested_frame[0]:interested_frame[1]].reset_index(drop=True)
  1378. }
  1379. location_value = get_lst_location_value(coord_data)
  1380. velocity_value = lst_data[m_id]['velocity']['waist'][interested_frame[0]:interested_frame[1]].reset_index(drop=True)
  1381. for zone in [0, 1, 2]:
  1382. zone_count = (location_value == zone).sum()
  1383. zone_pct = zone_count / len(location_value) * 100
  1384. zones_time_pcts[g][m][period_name][zone].append(zone_pct)
  1385. zone_speed = velocity_value[location_value == zone].mean() / pixpercm * framerate
  1386. zones_speed[g][m][period_name][zone].append(zone_speed)
  1387. # %%
  1388. # Convert zones_time_pcts to DataFrame (long format)
  1389. # Each row represents one mouse, first all Dom, then all Sub
  1390. import pandas as pd
  1391. import numpy as np
  1392. # Define order of time periods and zones
  1393. time_period_order = ['baseline', '0-5s', '5-20s', '20-60s']
  1394. zone_order = ['nest', 'edge', 'center']
  1395. df_zones_pct = {'s': None, 'p': None}
  1396. for g in ['s', 'p']:
  1397. data_list = []
  1398. for m in ['D', 'S']: # D (Dom) first, then S (Sub)
  1399. rank_name = 'Dom' if m == 'D' else 'Sub'
  1400. m_ids = mD_ids if m == 'D' else mS_ids
  1401. for m_id in m_ids:
  1402. # Extract B ID (e.g., 'B1M1(S)' -> 'B1')
  1403. b_id = m_id.split('M')[0]
  1404. # Initialize data row for this mouse
  1405. row_data = {
  1406. 'B_id': b_id,
  1407. 'rank': rank_name,
  1408. 'mouse_id': m_id,
  1409. }
  1410. # Iterate through each time period and zone
  1411. for period in time_period_order:
  1412. for zone_id, zone_name in enumerate(zone_order):
  1413. # Column name format: zone_period (e.g., nest_baseline)
  1414. col_name = f"{zone_name}_{period}"
  1415. # Get all data for this mouse in this time period and zone
  1416. values = zones_time_pcts[g][m][period][zone_id]
  1417. if period == 'baseline':
  1418. # Baseline: each mouse has only one value
  1419. # Find the index of this mouse in m_ids
  1420. mouse_idx = m_ids.index(m_id)
  1421. if mouse_idx < len(values):
  1422. row_data[col_name] = values[mouse_idx]
  1423. else:
  1424. row_data[col_name] = np.nan
  1425. else:
  1426. # Other time periods: calculate average across all looming events for this mouse
  1427. m_ids_list = mD_ids if m == 'D' else mS_ids
  1428. mouse_idx = m_ids_list.index(m_id)
  1429. # Calculate the number of looming events for this mouse
  1430. looming_count = len(looming_start_frames.get(m_id, []))
  1431. # Starting index for this mouse's data
  1432. start_idx = sum([len(looming_start_frames.get(mid, []))
  1433. for mid in m_ids_list[:mouse_idx]])
  1434. end_idx = start_idx + looming_count
  1435. # Get all looming data for this mouse
  1436. mouse_values = values[start_idx:end_idx]
  1437. # Calculate average
  1438. if len(mouse_values) > 0:
  1439. row_data[col_name] = np.nanmean(mouse_values)
  1440. else:
  1441. row_data[col_name] = np.nan
  1442. # Add to list
  1443. data_list.append(row_data)
  1444. # Convert to DataFrame
  1445. df_zones_pct[g] = pd.DataFrame(data_list)
  1446. # Adjust column order
  1447. # Information columns
  1448. info_cols = ['B_id', 'rank', 'mouse_id']
  1449. # Data columns
  1450. data_cols = [f"{zone}_{period}" for period in time_period_order for zone in zone_order]
  1451. for g in ['s', 'p']:
  1452. df_zones_pct[g] = df_zones_pct[g][info_cols + data_cols]
  1453. # Save DataFrame
  1454. df_zones_pct['s'].to_csv('data/FigS1B_zones_time_pcts.csv', index=False)
  1455. # %%
  1456. df = pd.read_csv('data/FigS1B_zones_time_pcts.csv')
  1457. print(df.columns.tolist())
  1458. exclude_cols = ['B_id', 'rank', 'mouse_id']
  1459. plot_cols = [col for col in df.columns if col not in exclude_cols]
  1460. means = df[plot_cols].mean()
  1461. plt.figure()
  1462. bar_colors = ['green', 'blue', 'orange']
  1463. bars = plt.bar(means.index, means.values, color=bar_colors[:len(means)])
  1464. df_melted = df[plot_cols]
  1465. sns.stripplot(
  1466. data=df_melted,
  1467. jitter=True,
  1468. color='gray',
  1469. size=5
  1470. )
  1471. plt.ylabel('Time in zone (%)')
  1472. plt.xticks(rotation=45, ha='right')
  1473. plt.show()
  1474. # %%
  1475. # Normalize data: convert from percentage to unit (s/(min∙cm²))
  1476. import pandas as pd
  1477. import numpy as np
  1478. # Define parameters
  1479. pixpergrid = 43.52
  1480. pixpercm = pixpergrid / 2 # pixels per cm
  1481. arena_size = 1088 # pixels
  1482. # Define time period durations (minutes)
  1483. time_durations_min = {
  1484. 'baseline': 150 / 60, # 2.5 min
  1485. '0-5s': 5 / 60, # 0.0833 min
  1486. '5-20s': 15 / 60, # 0.25 min
  1487. '20-60s': 40 / 60 # 0.667 min
  1488. }
  1489. # Define zone areas (cm²)
  1490. zone_areas_cm2 = {
  1491. 'nest': 80*4,
  1492. 'edge': 280*4,
  1493. 'center': 265*4
  1494. }
  1495. # Copy df_zones_wide
  1496. df_zones_unit = df_zones_pct['s'].copy()
  1497. # Normalize each data column
  1498. for col in data_cols:
  1499. # Parse column name: zone_period
  1500. parts = col.split('_')
  1501. if len(parts) >= 2:
  1502. zone = parts[0] # nest, edge, center
  1503. period = '_'.join(parts[1:]) # baseline, 0-5s, 5-20s, 20-60s
  1504. # Get time duration and zone area
  1505. time_min = time_durations_min.get(period, 1)
  1506. area_cm2 = zone_areas_cm2.get(zone, 1)
  1507. # Normalize: percentage / 100 / time(min) / area(cm²)
  1508. # Multiply by 60 because the result unit is s/(min∙cm²),
  1509. # converting percentage to decimal represents the proportion,
  1510. # proportion * total time(s) = actual time(s)
  1511. # actual time(s) / time(min) / area(cm²) = s/(min∙cm²)
  1512. df_zones_unit[col] = df_zones_unit[col] / 100 * 60 / area_cm2
  1513. # Save normalized data
  1514. df_zones_unit.to_csv('data/Fig1C_zones_time_unit.csv', index=False)
  1515. # %%
  1516. df = pd.read_csv('data/Fig1C_zones_time_unit.csv')
  1517. print(df.columns.tolist())
  1518. exclude_cols = ['B_id', 'rank', 'mouse_id']
  1519. plot_cols = [col for col in df.columns if col not in exclude_cols]
  1520. means = df[plot_cols].mean()
  1521. plt.figure()
  1522. bar_colors = ['green', 'blue', 'orange']
  1523. bars = plt.bar(means.index, means.values, color=bar_colors[:len(means)])
  1524. df_melted = df[plot_cols]
  1525. sns.stripplot(
  1526. data=df_melted,
  1527. jitter=True,
  1528. color='gray',
  1529. size=5
  1530. )
  1531. plt.ylabel('Time in zone (s/(min﹡cm²))')
  1532. plt.xticks(rotation=45, ha='right')
  1533. plt.show()
  1534. # %%
  1535. # Convert zones_speed to DataFrame (long format)
  1536. # Each row represents one mouse, first all Dom, then all Sub
  1537. import pandas as pd
  1538. import numpy as np
  1539. # Define order of time periods and zones
  1540. time_period_order = ['baseline', '0-5s', '5-20s', '20-60s']
  1541. zone_order = ['nest', 'edge', 'center']
  1542. df_zones_speed = {'s': None, 'p': None}
  1543. for g in ['s', 'p']:
  1544. data_list = []
  1545. for m in ['D', 'S']: # D (Dom) first, then S (Sub)
  1546. rank_name = 'Dom' if m == 'D' else 'Sub'
  1547. m_ids = mD_ids if m == 'D' else mS_ids
  1548. for m_id in m_ids:
  1549. # Extract B ID (e.g., 'B1M1(S)' -> 'B1')
  1550. b_id = m_id.split('M')[0]
  1551. # Initialize data row for this mouse
  1552. row_data = {
  1553. 'B_id': b_id,
  1554. 'rank': rank_name,
  1555. 'mouse_id': m_id,
  1556. }
  1557. # Iterate through each time period and zone
  1558. for period in time_period_order:
  1559. for zone_id, zone_name in enumerate(zone_order):
  1560. # Column name format: zone_period (e.g., nest_baseline)
  1561. col_name = f"{zone_name}_{period}"
  1562. # Get all data for this mouse in this time period and zone
  1563. values = zones_speed[g][m][period][zone_id]
  1564. if period == 'baseline':
  1565. # Baseline: each mouse has only one value
  1566. # Find the index of this mouse in m_ids
  1567. mouse_idx = m_ids.index(m_id)
  1568. if mouse_idx < len(values):
  1569. row_data[col_name] = values[mouse_idx]
  1570. else:
  1571. row_data[col_name] = np.nan
  1572. else:
  1573. # Other time periods: calculate average across all looming events for this mouse
  1574. m_ids_list = mD_ids if m == 'D' else mS_ids
  1575. mouse_idx = m_ids_list.index(m_id)
  1576. # Calculate the number of looming events for this mouse
  1577. looming_count = len(looming_start_frames.get(m_id, []))
  1578. # Starting index for this mouse's data
  1579. start_idx = sum([len(looming_start_frames.get(mid, []))
  1580. for mid in m_ids_list[:mouse_idx]])
  1581. end_idx = start_idx + looming_count
  1582. # Get all looming data for this mouse
  1583. mouse_values = values[start_idx:end_idx]
  1584. # Calculate average
  1585. if len(mouse_values) > 0:
  1586. row_data[col_name] = np.nanmean(mouse_values)
  1587. else:
  1588. row_data[col_name] = np.nan
  1589. # Add to list
  1590. data_list.append(row_data)
  1591. # Convert to DataFrame
  1592. df_zones_speed[g] = pd.DataFrame(data_list)
  1593. # Adjust column order
  1594. # Information columns
  1595. info_cols = ['B_id', 'rank', 'mouse_id']
  1596. # Data columns
  1597. data_cols = [f"{zone}_{period}" for period in time_period_order for zone in zone_order]
  1598. for g in ['s', 'p']:
  1599. df_zones_speed[g] = df_zones_speed[g][info_cols + data_cols]
  1600. # Save DataFrame
  1601. df_zones_speed['s'].to_csv('data/Fig1E_zones_speed.csv', index=False)
  1602. # %%
  1603. df = pd.read_csv('data/Fig1E_zones_speed.csv')
  1604. print(df.columns.tolist())
  1605. exclude_cols = ['B_id', 'rank', 'mouse_id']
  1606. plot_cols = [col for col in df.columns if col not in exclude_cols]
  1607. means = df[plot_cols].mean()
  1608. plt.figure()
  1609. bar_colors = ['green', 'blue', 'orange']
  1610. bars = plt.bar(means.index, means.values, color=bar_colors[:len(means)])
  1611. df_melted = df[plot_cols]
  1612. sns.stripplot(
  1613. data=df_melted,
  1614. jitter=True,
  1615. color='gray',
  1616. size=5
  1617. )
  1618. plt.ylabel('Speed in zone (cm/s)')
  1619. plt.xticks(rotation=45, ha='right')
  1620. plt.show()
  1621. # %% [markdown]
  1622. # ### Fig1F. ΔTime in zone (%)
  1623. # %%
  1624. import pandas as pd
  1625. import numpy as np
  1626. # Define order of time periods and zones
  1627. time_period_order = ['baseline', '0-5s', '5-20s', '20-60s']
  1628. zone_order = ['nest', 'edge', 'center']
  1629. for period in time_period_order:
  1630. period_cols = [col for col in df_zones_pct['s'].columns if period in col]
  1631. # Extract single and pair baseline data
  1632. df_single = df_zones_pct['s'][['B_id', 'rank'] + period_cols].copy()
  1633. df_pair = df_zones_pct['p'][['B_id', 'rank'] + period_cols].copy()
  1634. # Merge data (by B_id and rank)
  1635. df_merged = df_single.merge(df_pair, on=['B_id', 'rank'], suffixes=('_s', '_p'))
  1636. # Calculate difference (pair - single)
  1637. delta_data = []
  1638. for b_id in df_merged['B_id'].unique():
  1639. row = {'B_id': b_id}
  1640. b_data = df_merged[df_merged['B_id'] == b_id]
  1641. for zone in zone_order:
  1642. # Dom data
  1643. dom = b_data[b_data['rank'] == 'Dom']
  1644. if len(dom) > 0:
  1645. col_s = f'{zone}_{period}_s'
  1646. col_p = f'{zone}_{period}_p'
  1647. row[f'{zone}_D'] = dom[col_p].values[0] - dom[col_s].values[0]
  1648. else:
  1649. row[f'{zone}_D'] = np.nan
  1650. # Sub data
  1651. sub = b_data[b_data['rank'] == 'Sub']
  1652. if len(sub) > 0:
  1653. col_s = f'{zone}_{period}_s'
  1654. col_p = f'{zone}_{period}_p'
  1655. row[f'{zone}_S'] = sub[col_p].values[0] - sub[col_s].values[0]
  1656. else:
  1657. row[f'{zone}_S'] = np.nan
  1658. delta_data.append(row)
  1659. # Create DataFrame
  1660. df_delta_time = pd.DataFrame(delta_data)
  1661. # Adjust column order
  1662. df_delta_time = df_delta_time[
  1663. ['B_id', 'nest_D', 'nest_S', 'edge_D', 'edge_S', 'center_D', 'center_S']
  1664. ]
  1665. # Save results
  1666. df_delta_time.to_csv(f'data/Fig1F_delta_zones_time_{period}.csv', index=False)
  1667. # %%
  1668. for period in time_period_order:
  1669. df = pd.read_csv(f'data/Fig1F_delta_zones_time_{period}.csv')
  1670. groups = ['nest', 'edge', 'center']
  1671. fig, ax = plt.subplots()
  1672. x_pos = {}
  1673. xticks = []
  1674. xticklabels = []
  1675. x = 0
  1676. colors = {'D': 'orange', 'S': 'blue'}
  1677. # 1. x
  1678. for g in groups:
  1679. for cond in ['D', 'S']:
  1680. if cond == 'D':
  1681. x_pos[f'{g}_{cond}_mean'] = x
  1682. x_pos[f'{g}_{cond}_pts'] = x + 0.8
  1683. xticks += [x, x + 0.8]
  1684. xticklabels += [f'{g}_{cond}_mean', f'{g}_{cond}_pts']
  1685. else:
  1686. x_pos[f'{g}_{cond}_pts'] = x
  1687. x_pos[f'{g}_{cond}_mean'] = x + 0.8
  1688. xticks += [x, x + 0.8]
  1689. xticklabels += [f'{g}_{cond}_pts', f'{g}_{cond}_mean']
  1690. x += 2.2
  1691. # 2. mean ± SD
  1692. for g in groups:
  1693. for cond in ['D', 'S']:
  1694. col = f'{g}_{cond}'
  1695. xpos = x_pos[f'{g}_{cond}_mean']
  1696. ax.errorbar(
  1697. xpos,
  1698. df[col].mean(),
  1699. yerr=df[col].std(),
  1700. fmt='o',
  1701. color=colors[cond],
  1702. capsize=4,
  1703. markersize=9,
  1704. elinewidth=1.5,
  1705. capthick=1.5,
  1706. zorder=3
  1707. )
  1708. # 3. raw points + pairing
  1709. for i in range(len(df)):
  1710. for g in groups:
  1711. D_col = f'{g}_D'
  1712. S_col = f'{g}_S'
  1713. x_d = x_pos[f'{g}_D_pts']
  1714. y_d = df.loc[i, D_col]
  1715. x_s = x_pos[f'{g}_S_pts']
  1716. y_s = df.loc[i, S_col]
  1717. ax.scatter(x_d, y_d, color=colors['D'])
  1718. ax.scatter(x_s, y_s, color=colors['S'])
  1719. ax.plot([x_d, x_s], [y_d, y_s],
  1720. color='gray', linewidth=1)
  1721. # 4. axis
  1722. ax.set_xticks(xticks)
  1723. ax.set_xticklabels(xticklabels, rotation=45, ha='right')
  1724. ax.set_ylabel('Δ Time in zone (%)')
  1725. ax.set_ylim([-100, 100])
  1726. ax.set_title(f'Δ Time in zone ({period})')
  1727. plt.tight_layout()
  1728. plt.show()
  1729. # %% [markdown]
  1730. # ### Fig1G. Δ Speed in zone (cm/s)
  1731. # %%
  1732. import pandas as pd
  1733. import numpy as np
  1734. # Define order of time periods and zones
  1735. time_period_order = ['baseline', '0-5s', '5-20s', '20-60s']
  1736. zone_order = ['nest', 'edge', 'center']
  1737. for period in time_period_order:
  1738. period_cols = [col for col in df_zones_speed['s'].columns if period in col]
  1739. # Extract single and pair baseline data
  1740. df_single = df_zones_speed['s'][['B_id', 'rank'] + period_cols].copy()
  1741. df_pair = df_zones_speed['p'][['B_id', 'rank'] + period_cols].copy()
  1742. # Merge data (by B_id and rank)
  1743. df_merged = df_single.merge(df_pair, on=['B_id', 'rank'], suffixes=('_s', '_p'))
  1744. # Calculate difference (pair - single)
  1745. delta_data = []
  1746. for b_id in df_merged['B_id'].unique():
  1747. row = {'B_id': b_id}
  1748. b_data = df_merged[df_merged['B_id'] == b_id]
  1749. for zone in zone_order:
  1750. # Dom data
  1751. dom = b_data[b_data['rank'] == 'Dom']
  1752. if len(dom) > 0:
  1753. col_s = f'{zone}_{period}_s'
  1754. col_p = f'{zone}_{period}_p'
  1755. row[f'{zone}_D'] = dom[col_p].values[0] - dom[col_s].values[0]
  1756. else:
  1757. row[f'{zone}_D'] = np.nan
  1758. # Sub data
  1759. sub = b_data[b_data['rank'] == 'Sub']
  1760. if len(sub) > 0:
  1761. col_s = f'{zone}_{period}_s'
  1762. col_p = f'{zone}_{period}_p'
  1763. row[f'{zone}_S'] = sub[col_p].values[0] - sub[col_s].values[0]
  1764. else:
  1765. row[f'{zone}_S'] = np.nan
  1766. delta_data.append(row)
  1767. # Create DataFrame
  1768. df_delta_speed = pd.DataFrame(delta_data)
  1769. # Adjust column order
  1770. df_delta_speed = df_delta_speed[
  1771. ['B_id', 'nest_D', 'nest_S', 'edge_D', 'edge_S', 'center_D', 'center_S']
  1772. ]
  1773. # Save results
  1774. df_delta_speed.to_csv(f'data/Fig1G_delta_zones_speed_{period}.csv', index=False)
  1775. # %%
  1776. for period in time_period_order:
  1777. df = pd.read_csv(f'data/Fig1G_delta_zones_speed_{period}.csv')
  1778. groups = ['nest', 'edge', 'center']
  1779. fig, ax = plt.subplots()
  1780. x_pos = {}
  1781. xticks = []
  1782. xticklabels = []
  1783. x = 0
  1784. colors = {'D': 'orange', 'S': 'blue'}
  1785. # 1. x
  1786. for g in groups:
  1787. for cond in ['D', 'S']:
  1788. if cond == 'D':
  1789. x_pos[f'{g}_{cond}_mean'] = x
  1790. x_pos[f'{g}_{cond}_pts'] = x + 0.8
  1791. xticks += [x, x + 0.8]
  1792. xticklabels += [f'{g}_{cond}_mean', f'{g}_{cond}_pts']
  1793. else:
  1794. x_pos[f'{g}_{cond}_pts'] = x
  1795. x_pos[f'{g}_{cond}_mean'] = x + 0.8
  1796. xticks += [x, x + 0.8]
  1797. xticklabels += [f'{g}_{cond}_pts', f'{g}_{cond}_mean']
  1798. x += 2.2
  1799. # 2. mean ± SD
  1800. for g in groups:
  1801. for cond in ['D', 'S']:
  1802. col = f'{g}_{cond}'
  1803. xpos = x_pos[f'{g}_{cond}_mean']
  1804. ax.errorbar(
  1805. xpos,
  1806. df[col].mean(),
  1807. yerr=df[col].std(),
  1808. fmt='o',
  1809. color=colors[cond],
  1810. capsize=4,
  1811. markersize=9,
  1812. elinewidth=1.5,
  1813. capthick=1.5,
  1814. zorder=3
  1815. )
  1816. # 3. raw points + pairing
  1817. for i in range(len(df)):
  1818. for g in groups:
  1819. D_col = f'{g}_D'
  1820. S_col = f'{g}_S'
  1821. x_d = x_pos[f'{g}_D_pts']
  1822. y_d = df.loc[i, D_col]
  1823. x_s = x_pos[f'{g}_S_pts']
  1824. y_s = df.loc[i, S_col]
  1825. ax.scatter(x_d, y_d, color=colors['D'])
  1826. ax.scatter(x_s, y_s, color=colors['S'])
  1827. ax.plot([x_d, x_s], [y_d, y_s],
  1828. color='gray', linewidth=1)
  1829. # 4. axis
  1830. ax.set_xticks(xticks)
  1831. ax.set_xticklabels(xticklabels, rotation=45, ha='right')
  1832. ax.set_ylabel('Δ Speed in zone (cm/s)')
  1833. ax.set_ylim([-40, 40])
  1834. ax.set_title(f'Δ Speed in zone ({period})')
  1835. plt.tight_layout()
  1836. plt.show()
  1837. # %% [markdown]
  1838. # ### Fig1H. post-looming ethogram
  1839. # %%
  1840. import numpy as np
  1841. import pandas as pd
  1842. import matplotlib.pyplot as plt
  1843. import matplotlib.colors as mcolors
  1844. import colorsys
  1845. from matplotlib.patches import FancyArrow
  1846. from analyze_data_utils import filter_in_range
  1847. def adjust_saturation(color, saturation):
  1848. rgb = mcolors.to_rgb(color)
  1849. h, l, s = colorsys.rgb_to_hls(*rgb)
  1850. new_s = saturation * s
  1851. new_rgb = colorsys.hls_to_rgb(h, l, new_s)
  1852. return new_rgb
  1853. with open('data/lst_etho_dict.pkl', 'rb') as f:
  1854. lst_etho_dict = pickle.load(f)
  1855. lst_framerange = [0, 1950]
  1856. saturation = 0.9
  1857. framerate=30
  1858. stim_frame = int(0.725 * framerate)
  1859. behavior_frames_dict = lst_etho_dict
  1860. behavior_frames_dict = dict(sorted(behavior_frames_dict.items()))
  1861. lst_behavior_frames_dict = filter_in_range(lst_etho_dict, lst_framerange, method='replace')
  1862. lst_behavior_properties = {
  1863. 'approach_partner': ('#143FCA', 4),
  1864. 'follow_partner': ('#137CAB', 4),
  1865. 'groom_partner': ('#0990FF', 4),
  1866. 'sniff_partner': ('#09FFFF', 4),
  1867. 'tailrattling': ('#FF09FF', 5),
  1868. 'huddling': ('#6D57F3', 4),
  1869. 'jumping': ('#FF1717', 5),
  1870. 'escape': ('#FF1717', 4),
  1871. 'freezing': ('#FF75EF', 4),
  1872. 'dwelling': ('#FF75EF', 4),
  1873. 'grooming': ('#28AE61', 4),
  1874. 'real_rearing':('#C409FF', 3),
  1875. 'stretching': ('#C409FF', 3),
  1876. 'nest': ('gray', 2),
  1877. 'sniffing': ('#0C8140', 1),
  1878. 'climbing': ('gray', 1),
  1879. }
  1880. fig, axes = plt.subplots(1, 2, figsize=(10, 2), dpi=300)
  1881. plt.subplots_adjust(wspace=0)
  1882. for i, group in enumerate(lst_etho_dict.keys()):
  1883. ax = axes[i]
  1884. ax.set_yticks(range(len(lst_etho_dict[group].keys())))
  1885. for idx, t_id in enumerate(lst_etho_dict[group].keys()):
  1886. behaviors = lst_etho_dict[group][t_id]
  1887. for b, frame_ranges in behaviors.items():
  1888. for frame_range in frame_ranges:
  1889. if np.isnan(frame_ranges).all():
  1890. continue
  1891. start_frame, end_frame = frame_range
  1892. color, zorder = lst_behavior_properties.get(b, ('white', 0))
  1893. color = adjust_saturation(color, saturation)
  1894. rect = plt.Rectangle((start_frame, idx - 0.4), end_frame - start_frame, 0.8, facecolor=color, edgecolor='none', zorder=zorder)
  1895. ax.add_patch(rect)
  1896. if idx % 2 == 0:
  1897. ax.axhline(idx - 0.5, color='black', linestyle='-', linewidth=1, zorder=9)
  1898. ytick_color = 'red'
  1899. else:
  1900. ax.axhline(idx - 0.5, color='black', linestyle='-', linewidth=0.1, zorder=9)
  1901. ytick_color = 'green'
  1902. ax.get_yticklabels()[idx].set_color(ytick_color)
  1903. span_start_frame = 5 * framerate
  1904. span_end_frame = span_start_frame + stim_frame
  1905. if i == 0:
  1906. ax.set_yticklabels(['D', 'S', 'D', 'S', 'D', 'S'])
  1907. ax.set_ylabel('Looming')
  1908. else:
  1909. ax.set_yticks([])
  1910. ax.set_xticks([])
  1911. lst_blank_width = (lst_framerange[1] - lst_framerange[0]) / 20
  1912. per_rect = plt.Rectangle((lst_framerange[0] - lst_blank_width, -0.5), lst_blank_width, 6,
  1913. facecolor='white', edgecolor='white', zorder=8)
  1914. ax.add_patch(per_rect)
  1915. post_rect = plt.Rectangle((lst_framerange[1], -0.5), lst_blank_width, 6,
  1916. facecolor='white', edgecolor='white', zorder=8)
  1917. ax.add_patch(post_rect)
  1918. ax.set_xlim(lst_framerange[0] - lst_blank_width, lst_framerange[1] + lst_blank_width)
  1919. ax.set_ylim(-0.5, 5.5)
  1920. ax.invert_yaxis()
  1921. ax.set_title(f'{group}')
  1922. lst_arrow = FancyArrow(30*5, -1.1, 0, 0.5, width=lst_blank_width/50, head_width=lst_blank_width/10, head_length=0.3, length_includes_head=True, color='red')
  1923. lst_arrow.set_clip_on(False)
  1924. ax.add_patch(lst_arrow)
  1925. threat_line = plt.Line2D([span_start_frame, span_end_frame], [5.75, 5.75], color='red', linewidth=2)
  1926. threat_line.set_clip_on(False)
  1927. ax.add_artist(threat_line)
  1928. if i == 1:
  1929. line_x1 = lst_framerange[1] - 30*10
  1930. line_x2 = lst_framerange[1]
  1931. ax.text((line_x1 + line_x2) / 2, -0.8, '10 sec', fontsize=8, color='black', ha='center')
  1932. lst_line = plt.Line2D([line_x1, line_x2], [-0.65, -0.65], color='black', linewidth=1)
  1933. lst_line.set_clip_on(False)
  1934. ax.add_artist(lst_line)
  1935. lst_legend_ = {
  1936. 'escape': ('#FF1717', 4),
  1937. 'freezing': ('#FF75EF', 4),
  1938. 'tail rattling': ('#FF09FF', 5),
  1939. 'stretching / rearing': ('#C409FF', 3),
  1940. 'huddling': ('#6D57F3', 1),
  1941. 'approaching P (partner)': ('#143FCA', 4),
  1942. 'following P': ('#137CAB', 4),
  1943. 'grooming P': ('#0990FF', 4),
  1944. 'sniffing P': ('#09FFFF', 4),
  1945. 'grooming': ('#28AE61', 4),
  1946. 'sniffing': ('#0C8140', 1),
  1947. 'in nest': ('gray', 2),
  1948. 'other behaviors': ('white', 1)}
  1949. lst_legend = []
  1950. for label, (color, zorder) in lst_legend_.items():
  1951. if label == 'other behaviors':
  1952. rect = plt.Rectangle((0, 0), 1, 1, facecolor=adjust_saturation(color, saturation),
  1953. edgecolor='black', linewidth=0.5, label=label)
  1954. else:
  1955. rect = plt.Rectangle((0, 0), 1, 1, facecolor=adjust_saturation(color, saturation), label=label)
  1956. lst_legend.append(rect)
  1957. axes[1].legend(handles=lst_legend, loc='upper center', bbox_to_anchor=(0, -0.2), fontsize=7, ncol=7)
  1958. plt.savefig('fig/fig_1/lst_behavior_ethogram.eps', format="eps", dpi=300, bbox_inches="tight")
  1959. plt.show()
  1960. # %% [markdown]
  1961. # ### Fig1I. Social time(s/(min∙cm2))
  1962. # %%
  1963. import copy
  1964. import numpy as np
  1965. import pandas as pd
  1966. from analyze_data_utils import dict_to_dataframe, filter_dict_data, merge_dicts, filter_in_range, calculate_duration_time, calculate_total
  1967. def filter_social_in_nest(social_frames, location_series):
  1968. """
  1969. social_frames: [(start, end), ...]
  1970. location_series: pd.Series, starting from 0, 0 indicates in nest
  1971. Returns: [(nest_start, nest_end), ...]
  1972. """
  1973. social_in_nest = []
  1974. for start, end in social_frames:
  1975. # Extract location values for this interval
  1976. loc_segment = location_series[start:end+1] # +1 ensures inclusion of end frame
  1977. # Find consecutive segments equal to 0
  1978. in_nest_mask = (loc_segment == 0)
  1979. if not in_nest_mask.any():
  1980. continue # This social segment is completely outside the nest
  1981. # Find consecutive True segments
  1982. in_nest_indices = loc_segment.index[in_nest_mask]
  1983. group_start = None
  1984. for idx in in_nest_indices:
  1985. if group_start is None:
  1986. group_start = idx
  1987. prev_idx = idx
  1988. elif idx == prev_idx + 1:
  1989. prev_idx = idx
  1990. else:
  1991. social_in_nest.append((group_start, prev_idx))
  1992. group_start = idx
  1993. prev_idx = idx
  1994. # Last segment
  1995. if group_start is not None:
  1996. social_in_nest.append((group_start, prev_idx))
  1997. if social_in_nest == []:
  1998. social_in_nest = [np.nan]
  1999. return social_in_nest
  2000. ap_bl_data = filter_dict_data(lst_baseline_bhvr_labels_dict, 'approach_partner')
  2001. sp_bl_data = filter_dict_data(lst_baseline_bhvr_labels_dict, 'sniff_partner')
  2002. fp_bl_data = filter_dict_data(lst_baseline_bhvr_labels_dict, 'follow_partner')
  2003. gp_bl_data = filter_dict_data(lst_baseline_bhvr_labels_dict, 'groom_partner')
  2004. hp_bl_data = filter_dict_data(lst_baseline_bhvr_labels_dict, 'huddling')
  2005. merged_bl_data = merge_dicts(ap_bl_data, sp_bl_data)
  2006. merged_bl_data = merge_dicts(merged_bl_data, fp_bl_data)
  2007. merged_bl_data = merge_dicts(merged_bl_data, gp_bl_data)
  2008. merged_bl_data = merge_dicts(merged_bl_data, hp_bl_data)
  2009. ap_wr_data = filter_dict_data(lst_bhvr_labels_dict, 'approach_partner')
  2010. sp_wr_data = filter_dict_data(lst_bhvr_labels_dict, 'sniff_partner')
  2011. fp_wr_data = filter_dict_data(lst_bhvr_labels_dict, 'follow_partner')
  2012. gp_wr_data = filter_dict_data(lst_bhvr_labels_dict, 'groom_partner')
  2013. hp_wr_data = filter_dict_data(lst_bhvr_labels_dict, 'huddling')
  2014. merged_wr_data = merge_dicts(ap_wr_data, sp_wr_data)
  2015. merged_wr_data = merge_dicts(merged_wr_data, fp_wr_data)
  2016. merged_wr_data = merge_dicts(merged_wr_data, gp_wr_data)
  2017. merged_wr_data = merge_dicts(merged_wr_data, hp_wr_data)
  2018. merged_bl_data = filter_in_range(merged_bl_data, [0,4500], method='replace')
  2019. # merged_wr_data = filter_in_range(merged_wr_data, [150,1950], method='replace')
  2020. # merged_wr_data = filter_in_range(merged_wr_data, [150,300], method='replace')
  2021. merged_wr_data = filter_in_range(merged_wr_data, [150,1950], method='replace')
  2022. merged_data = {'bl': merged_bl_data['p'], 'pl': merged_wr_data['p']}
  2023. # merged_data_ = merge_intervals(merged_data, max_gap=30)
  2024. merged_data_ = merged_data
  2025. g = 'p'
  2026. social_in_nest_data = {}
  2027. for m in ['D','S']:
  2028. for t in ['bl', 'pl']:
  2029. m_ids = mS_ids if m =='S' else mD_ids
  2030. lst_data = lsts_data if g == 's' else lstp_data
  2031. for m_id in m_ids:
  2032. mp_id = [k for k in mp_ids if m_id in k]
  2033. if t == 'bl':
  2034. interested_frame = (9000, 9000+4500)
  2035. t_id = m_id.split('M')[0]+'M'+m_id.split('(')[1][0]+'T1'
  2036. coord_data = {'x': lst_data[m_id]['coordinate']['waist']['x'][interested_frame[0]:interested_frame[1]].reset_index(drop=True),
  2037. 'y': lst_data[m_id]['coordinate']['waist']['y'][interested_frame[0]:interested_frame[1]].reset_index(drop=True)}
  2038. location_value = get_lst_location_value(coord_data)
  2039. social_frames = merged_data_[t][m][t_id]
  2040. if t not in social_in_nest_data:
  2041. social_in_nest_data[t] = {}
  2042. if m not in social_in_nest_data[t]:
  2043. social_in_nest_data[t][m] = {}
  2044. if not np.isnan(social_frames).any():
  2045. social_in_nest = filter_social_in_nest(social_frames, location_value)
  2046. social_in_nest_data[t][m][t_id] = social_in_nest
  2047. else:
  2048. social_in_nest_data[t][m][t_id] = [np.nan]
  2049. elif t == 'pl':
  2050. for idx, lsf in enumerate(looming_start_frames[mp_id[0]], start=1):
  2051. interested_frame = (lsf+30*0, lsf+30*60)
  2052. t_id = m_id.split('M')[0]+'M'+m_id.split('(')[1][0]+'T'+str(idx)
  2053. # interes20ted_frame = (looming_start_frames[m_id][0], looming_start_frames[m_id][-1]+30*60)
  2054. # for lsf in looming_start_frames[m_id]:
  2055. # t_id = m_id.split('M')[0]+'M'+m_id.split('(')[1][0]+'T1'
  2056. # for nest in nest_data[g][m][t_id]:
  2057. # if not np.isnan(nest).any():
  2058. # start, end = nest
  2059. # interested_frame = (9000+start-60, 9000+start+10)
  2060. # interested_frame = (lsf-150+start-60, lsf-150+start+10)
  2061. coord_data = {'x': lst_data[m_id]['coordinate']['waist']['x'][interested_frame[0]:interested_frame[1]].reset_index(drop=True),
  2062. 'y': lst_data[m_id]['coordinate']['waist']['y'][interested_frame[0]:interested_frame[1]].reset_index(drop=True)}
  2063. location_value = get_lst_location_value(coord_data)
  2064. if t_id in merged_data_[t][m].keys():
  2065. social_frames = merged_data_[t][m][t_id]
  2066. if t not in social_in_nest_data:
  2067. social_in_nest_data[t] = {}
  2068. if m not in social_in_nest_data[t]:
  2069. social_in_nest_data[t][m] = {}
  2070. if not np.isnan(social_frames).any():
  2071. social_in_nest = filter_social_in_nest(social_frames, location_value)
  2072. social_in_nest_data[t][m][t_id] = social_in_nest
  2073. else:
  2074. social_in_nest_data[t][m][t_id] = [np.nan]
  2075. raw_social_in_nest_data = copy.deepcopy(social_in_nest_data)
  2076. social_in_nest_data = {'bl': raw_social_in_nest_data['bl'], 'pl': raw_social_in_nest_data['pl']}
  2077. social_in_nest_durations = calculate_duration_time(social_in_nest_data)
  2078. total_social_in_nest_duration = calculate_total(social_in_nest_durations)
  2079. total_social_in_nest_duration_df = dict_to_dataframe(total_social_in_nest_duration, value_name='in_nest_duration', groups=['bl', 'pl'], nan2zero=True)
  2080. # total_social_in_nest_duration_df.to_csv('data/lst_total_social_in_nest_duration.csv')
  2081. raw_social_data = copy.deepcopy(merged_data_)
  2082. social_data = {'bl': raw_social_data['bl'], 'pl': raw_social_data['pl']}
  2083. social_durations = calculate_duration_time(social_data)
  2084. total_social_duration = calculate_total(social_durations)
  2085. total_social_duration_df = dict_to_dataframe(total_social_duration, value_name='total_social_duration', groups=['bl', 'pl'], nan2zero=True)
  2086. # Create total_social_out_nest_duration_df, calculate difference for each suffixed column separately
  2087. total_social_out_nest_duration_df = total_social_duration_df[['B_id']].copy()
  2088. # Calculate difference for each suffix (_blD, _blS, _plD, _plS) separately
  2089. for suffix in ['_blD', '_blS', '_plD', '_plS']:
  2090. total_col = f'total_social_duration{suffix}'
  2091. in_nest_col = f'in_nest_duration{suffix}'
  2092. out_nest_col = f'out_nest_duration{suffix}'
  2093. if total_col in total_social_duration_df.columns and in_nest_col in total_social_in_nest_duration_df.columns:
  2094. total_social_out_nest_duration_df[out_nest_col] = (
  2095. total_social_duration_df[total_col] - total_social_in_nest_duration_df[in_nest_col]
  2096. )
  2097. # total_social_out_nest_duration_df.to_csv('data/lst_total_social_out_nest_duration.csv')
  2098. # Merge in_nest and out_nest data, in first, out later
  2099. # Define parameters
  2100. baseline_time = 2.5 # 2.5 min = 150 seconds
  2101. afterlooming_time = 1 # 1 min = 60 seconds
  2102. in_nest_area = 300 # cm²
  2103. out_nest_area = 2500 - 300 # 2200 cm²
  2104. # Merge DataFrame
  2105. social_density_df = total_social_in_nest_duration_df[['B_id']].copy()
  2106. # Add columns in order: in_nest first, then out_nest
  2107. for condition in ['_blD', '_blS', '_plD', '_plS']:
  2108. # Add in_nest column first
  2109. in_col = f'in_nest_duration{condition}'
  2110. if in_col in total_social_in_nest_duration_df.columns:
  2111. # Determine if baseline (bl) or after-looming (pl)
  2112. time_divisor = baseline_time if '_bl' in condition else afterlooming_time
  2113. # Calculate density: (duration / time) / area
  2114. social_density_df[in_col] = (
  2115. total_social_in_nest_duration_df[in_col] / time_divisor / in_nest_area
  2116. )
  2117. # Add out_nest column later
  2118. out_col = f'out_nest_duration{condition}'
  2119. if out_col in total_social_out_nest_duration_df.columns:
  2120. # Determine if baseline (bl) or after-looming (pl)
  2121. time_divisor = baseline_time if '_bl' in condition else afterlooming_time
  2122. # Calculate density: (duration / time) / area
  2123. social_density_df[out_col] = (
  2124. total_social_out_nest_duration_df[out_col] / time_divisor / out_nest_area
  2125. )
  2126. social_density_df.to_csv('data/Fig1I_social_time_density.csv', index=False)
  2127. # %%
  2128. df = pd.read_csv('data/Fig1I_social_time_density.csv')
  2129. ordered_cols = ['in_nest_duration_blD', 'out_nest_duration_blD', 'in_nest_duration_blS', 'out_nest_duration_blS', 'in_nest_duration_plD', 'out_nest_duration_plD', 'in_nest_duration_plS', 'out_nest_duration_plS']
  2130. fig, ax = plt.subplots()
  2131. x_pos = {}
  2132. xticks = []
  2133. xticklabels = []
  2134. x = 0
  2135. colors = {'D': 'orange', 'S': 'blue'}
  2136. for col in ordered_cols:
  2137. cond = col[-1]
  2138. is_out = 'out' in col
  2139. if not is_out:
  2140. x_pos[f'{col}_mean'] = x
  2141. x_pos[f'{col}_pts'] = x + 0.8
  2142. xticks += [x, x + 0.8]
  2143. xticklabels += [f'{col}_mean', f'{col}_pts']
  2144. else:
  2145. x_pos[f'{col}_pts'] = x
  2146. x_pos[f'{col}_mean'] = x + 0.8
  2147. xticks += [x, x + 0.8]
  2148. xticklabels += [f'{col}_pts', f'{col}_mean']
  2149. x += 2.2
  2150. for col in ordered_cols:
  2151. cond = col[-1]
  2152. ax.errorbar(x_pos[f'{col}_mean'], df[col].mean(), yerr=df[col].std(), fmt='o', color=colors[cond], capsize=4, markersize=9, elinewidth=1.5, capthick=1.5, zorder=3)
  2153. for i in range(len(df)):
  2154. for c in ['bl', 'pl']:
  2155. for cond in ['D', 'S']:
  2156. in_col = f'in_nest_duration_{c}{cond}'
  2157. out_col = f'out_nest_duration_{c}{cond}'
  2158. x_in = x_pos[f'{in_col}_pts']
  2159. x_out = x_pos[f'{out_col}_pts']
  2160. y_in = df.loc[i, in_col]
  2161. y_out = df.loc[i, out_col]
  2162. ax.plot([x_in, x_out], [y_in, y_out], color='gray', linewidth=1)
  2163. ax.scatter(x_in, y_in, color=colors[cond])
  2164. ax.scatter(x_out, y_out, color=colors[cond])
  2165. ax.set_xticks(xticks)
  2166. ax.set_xticklabels(xticklabels, rotation=45, ha='right')
  2167. ax.set_ylabel('Social time (s/(min﹡cm²))')
  2168. plt.tight_layout()
  2169. plt.show()
  2170. # %% [markdown]
  2171. # ### Fig1K. Allocation of behavioral decisions (%)
  2172. # %%
  2173. import numpy as np
  2174. import pandas as pd
  2175. combine = True
  2176. row_names = ['E', 'F+E', 'F', 'I']
  2177. colors = ['red', 'purple', 'blue', 'green']
  2178. exclude = ['N', 'n']
  2179. # exclude = ['N', 'n', 'I', 'A']
  2180. def classify_behaviors(behaviors_dict, key1, key2, all_behaviors, exclude=None, combine=False):
  2181. behaviors_classified = []
  2182. for k1 in key1:
  2183. for k2 in key2:
  2184. behaviors_classified.extend(list(behaviors_dict[k1][k2].values()))
  2185. if exclude:
  2186. behaviors_classified = [behavior for behavior in behaviors_classified if all(excl not in behavior for excl in exclude)]
  2187. if combine:
  2188. behaviors_classified = ['E' if behavior == 'A+E' else 'I' if behavior == 'A' else behavior for behavior in behaviors_classified]
  2189. behaviors, indices = np.unique(behaviors_classified, return_inverse=True)
  2190. counts = np.bincount(indices)
  2191. percents = np.bincount(indices) / len(behaviors_classified) * 100
  2192. n = len(behaviors_classified)
  2193. behavior_counts = {behavior: 0 for behavior in all_behaviors}
  2194. behavior_percents = {behavior: 0.0 for behavior in all_behaviors}
  2195. for behavior, count, percent in zip(behaviors, counts, percents):
  2196. behavior_counts[behavior] = count
  2197. behavior_percents[behavior] = percent
  2198. sorted_counts = [behavior_counts[behavior] for behavior in all_behaviors]
  2199. sorted_percents = [behavior_percents[behavior] for behavior in all_behaviors]
  2200. return sorted_percents, sorted_counts, n
  2201. def calculate_behavior_decision_percentages(bhvr_deci_dict, behavior_names, exclude_list=None, combine_flag=False):
  2202. sD_deci_pct, sD_deci_cnt, sD_deci_n = classify_behaviors(
  2203. bhvr_deci_dict, ['s'], ['D'], behavior_names, exclude=exclude_list, combine=combine_flag)
  2204. pD_deci_pct, pD_deci_cnt, pD_deci_n = classify_behaviors(
  2205. bhvr_deci_dict, ['p'], ['D'], behavior_names, exclude=exclude_list, combine=combine_flag)
  2206. sS_deci_pct, sS_deci_cnt, sS_deci_n = classify_behaviors(
  2207. bhvr_deci_dict, ['s'], ['S'], behavior_names, exclude=exclude_list, combine=combine_flag)
  2208. pS_deci_pct, pS_deci_cnt, pS_deci_n = classify_behaviors(
  2209. bhvr_deci_dict, ['p'], ['S'], behavior_names, exclude=exclude_list, combine=combine_flag)
  2210. pct_df = pd.DataFrame({
  2211. 'behavior': behavior_names,
  2212. 'single_Dom': sD_deci_pct,
  2213. 'single_Sub': sS_deci_pct,
  2214. 'pair_Dom': pD_deci_pct,
  2215. 'pair_Sub': pS_deci_pct
  2216. })
  2217. cnt_df = pd.DataFrame({
  2218. 'behavior': behavior_names,
  2219. 'single_Dom': sD_deci_cnt,
  2220. 'single_Sub': sS_deci_cnt,
  2221. 'pair_Dom': pD_deci_cnt,
  2222. 'pair_Sub': pS_deci_cnt
  2223. })
  2224. n_df = pd.DataFrame({
  2225. 'group_rank': ['single_Dom', 'single_Sub', 'pair_Dom', 'pair_Sub'],
  2226. 'n': [sD_deci_n, sS_deci_n, pD_deci_n, pS_deci_n]
  2227. })
  2228. return pct_df, cnt_df, n_df
  2229. pct_df, cnt_df, n_df = calculate_behavior_decision_percentages(
  2230. lst_deci_dict, row_names, exclude_list=exclude, combine_flag=combine)
  2231. print(cnt_df)
  2232. pct_df.to_csv('data/Fig1K_lst_behavior_decision_pct.csv', index=False)
  2233. cnt_df.to_csv('data/Fig1K_lst_behavior_decision_cnt.csv', index=False)
  2234. # %%
  2235. import pandas as pd
  2236. import numpy as np
  2237. import matplotlib.pyplot as plt
  2238. df = pd.read_csv('data/Fig1K_lst_behavior_decision_pct.csv')
  2239. colors = {'E': 'red', 'F+E': 'purple', 'F': 'blue', 'I': 'green'}
  2240. x_labels = ['single_Dom', 'pair_Dom', 'single_Sub', 'pair_Sub']
  2241. x = np.arange(len(x_labels))
  2242. width = 0.6
  2243. fig, ax = plt.subplots()
  2244. bottom = np.zeros(len(x_labels))
  2245. for i, row in df.iloc[::-1].iterrows(): # reverse
  2246. vals = [row['single_Dom'], row['pair_Dom'], row['single_Sub'], row['pair_Sub']]
  2247. c = colors[row['behavior']]
  2248. ax.bar(x, vals, width, bottom=bottom, color=c)
  2249. bottom += vals
  2250. ax.set_xticks(x)
  2251. ax.set_xticklabels(x_labels)
  2252. ax.set_ylabel('Allocation of bahavioral decisitons (%)')
  2253. ax.set_ylim(0, 100)
  2254. plt.tight_layout()
  2255. plt.show()
  2256. # %% [markdown]
  2257. # ### Fig1L. Partner-self behavior combinations
  2258. # %%
  2259. import pandas as pd
  2260. import numpy as np
  2261. # Step 1: Merge all behaviors of a trial into a single categorical series
  2262. def merge_lst_behaviors_to_categories(trial_data, frame_range):
  2263. start_frame, end_frame = frame_range
  2264. categorized = pd.Series(['O'] * (end_frame - start_frame),
  2265. index=range(start_frame, end_frame))
  2266. escape_behaviors = ['escape']
  2267. freezing_behaviors = ['freezing', 'tail_rattling']
  2268. social_behaviors = ['approach_partner', 'sniff_partner', 'follow_partner', 'groom_partner', 'huddling']
  2269. for behavior in social_behaviors:
  2270. if behavior in trial_data:
  2271. behavior_series = trial_data[behavior]
  2272. for frame_idx in categorized.index:
  2273. if frame_idx in behavior_series.index and behavior_series[frame_idx] == 1:
  2274. categorized[frame_idx] = 'S'
  2275. for behavior in freezing_behaviors:
  2276. if behavior in trial_data:
  2277. behavior_series = trial_data[behavior]
  2278. for frame_idx in categorized.index:
  2279. if frame_idx in behavior_series.index and behavior_series[frame_idx] == 1:
  2280. categorized[frame_idx] = 'F'
  2281. for behavior in escape_behaviors:
  2282. if behavior in trial_data:
  2283. behavior_series = trial_data[behavior]
  2284. for frame_idx in categorized.index:
  2285. if frame_idx in behavior_series.index and behavior_series[frame_idx] == 1:
  2286. categorized[frame_idx] = 'E'
  2287. return categorized
  2288. # Step 2: Pair partner and self behavior series to generate combination series
  2289. def create_partner_self_combinations(self_series, partner_series):
  2290. """
  2291. Pair partner and self behavior series to produce combination labels.
  2292. Parameters:
  2293. self_series: pandas Series, self behavior classification
  2294. partner_series: pandas Series, partner behavior classification
  2295. Returns:
  2296. pandas Series with combination labels like 'E&F', 'F&F', 'O&O', etc.
  2297. """
  2298. if len(self_series) != len(partner_series):
  2299. raise ValueError("self_series and partner_series length mismatch")
  2300. combination_series = pd.Series(
  2301. [f"{partner_series.iloc[i]}&{self_series.iloc[i]}"
  2302. for i in range(len(self_series))],
  2303. index=self_series.index
  2304. )
  2305. return combination_series
  2306. # Step 3: Calculate time percentage for each combination
  2307. def calculate_combination_percentages(combination_series, all_combinations):
  2308. """
  2309. Compute percentage of frames for each combination.
  2310. Parameters:
  2311. combination_series: pandas Series with combination labels
  2312. Returns:
  2313. dict with percentages for all combinations
  2314. """
  2315. total_frames = len(combination_series)
  2316. percentages = {}
  2317. for combo in all_combinations:
  2318. count = (combination_series == combo).sum()
  2319. percentages[combo] = (count / total_frames * 100) if total_frames > 0 else 0.0
  2320. return percentages
  2321. # Step 4: Helper function to get partner session ID
  2322. def get_partner_session_id(session_id):
  2323. """
  2324. Derive partner session ID from a given session ID.
  2325. Example: B1MDT1 -> B1MST1, B1MST1 -> B1MDT1
  2326. """
  2327. parts = session_id.split('M')
  2328. if len(parts) != 2:
  2329. return None
  2330. base = parts[0] # e.g. "B1"
  2331. rest = parts[1] # e.g. "DT1" or "ST1"
  2332. if rest.startswith('D'):
  2333. partner_rest = 'S' + rest[1:] # D -> S
  2334. elif rest.startswith('S'):
  2335. partner_rest = 'D' + rest[1:] # S -> D
  2336. else:
  2337. return None
  2338. return base + 'M' + partner_rest
  2339. # Step 5: Main processing – analyse all pair trials
  2340. frame_start = 150
  2341. frame_end = 300
  2342. frame_range = (frame_start, frame_end)
  2343. lst_all_combs = ['E&E', 'E&F', 'E&S', 'E&O', 'F&E', 'F&F', 'F&S', 'F&O',
  2344. 'S&E', 'S&F', 'S&S', 'S&O', 'O&E', 'O&F', 'O&S', 'O&O']
  2345. # Initialize result dictionary
  2346. lst_comb_bhvr_dict = {'p': {'D': {}, 'S': {}}}
  2347. # Iterate over all pair group trials
  2348. group = 'p'
  2349. if group in lst_bhvr_frames_dict:
  2350. for rank in ['D', 'S']:
  2351. if rank not in lst_bhvr_frames_dict[group]:
  2352. continue
  2353. for session_id in sorted(lst_bhvr_frames_dict[group][rank].keys()):
  2354. # Get self behavior data
  2355. self_trial_data = lst_bhvr_frames_dict[group][rank][session_id]
  2356. # Determine partner rank and session ID
  2357. partner_rank = 'S' if rank == 'D' else 'D'
  2358. partner_session_id = get_partner_session_id(session_id)
  2359. if partner_session_id is None:
  2360. continue
  2361. # Check if partner data exists
  2362. if (partner_rank not in lst_bhvr_frames_dict[group] or
  2363. partner_session_id not in lst_bhvr_frames_dict[group][partner_rank]):
  2364. continue
  2365. partner_trial_data = lst_bhvr_frames_dict[group][partner_rank][partner_session_id]
  2366. # Step 1: Categorize self and partner behaviors separately
  2367. self_categorized = merge_lst_behaviors_to_categories(self_trial_data, frame_range)
  2368. partner_categorized = merge_lst_behaviors_to_categories(partner_trial_data, frame_range)
  2369. # Step 2: Create partner‑self combination series
  2370. combination_series = create_partner_self_combinations(
  2371. self_categorized, partner_categorized
  2372. )
  2373. # Step 3: Compute percentages for all 16 combinations
  2374. percentages = calculate_combination_percentages(combination_series, lst_all_combs)
  2375. # Step 4: Store in result dictionary
  2376. lst_comb_bhvr_dict['p'][rank][session_id] = percentages
  2377. # %%
  2378. from analyze_data_utils import filter_dict_data, dict_to_dataframe
  2379. import matplotlib.pyplot as plt
  2380. import seaborn as sns
  2381. import numpy as np
  2382. # Store DataFrames for all combinations
  2383. all_dfs = {}
  2384. # Process each combination
  2385. for comb in lst_all_combs:
  2386. # Filter data for this combination
  2387. lst_comb_bhvr_dict_filtered = filter_dict_data(lst_comb_bhvr_dict, comb)
  2388. # Convert to DataFrame
  2389. lst_comb_bhvr_df = dict_to_dataframe(
  2390. lst_comb_bhvr_dict_filtered,
  2391. groups=['p'],
  2392. ranks=['D', 'S'],
  2393. value_name=f'{comb}_pct',
  2394. nan2zero=True
  2395. )
  2396. all_dfs[comb] = lst_comb_bhvr_df
  2397. # Merge all combination data and compute average for dominant mice
  2398. dom_avg_data = {}
  2399. for comb in lst_all_combs:
  2400. df = all_dfs[comb]
  2401. # Select columns for dominant mice (column names ending with '_pD')
  2402. dom_cols = [col for col in df.columns if col.endswith('_pD')]
  2403. if len(dom_cols) > 0:
  2404. dom_values = df[dom_cols].values.flatten()
  2405. dom_avg_data[comb] = dom_values.mean() if len(dom_values) > 0 else 0.0
  2406. else:
  2407. dom_avg_data[comb] = 0.0
  2408. # Build a 4x4 matrix from dom_avg_data
  2409. categories = ['E', 'F', 'S', 'O']
  2410. matrix_4x4 = np.zeros((4, 4))
  2411. for i, cat1 in enumerate(categories):
  2412. for j, cat2 in enumerate(categories):
  2413. comb = f"{cat1}&{cat2}"
  2414. if comb in dom_avg_data:
  2415. matrix_4x4[i, j] = dom_avg_data[comb]
  2416. # Create 4x4 DataFrame
  2417. lst_dom_bhvr_comb_4x4 = pd.DataFrame(
  2418. matrix_4x4,
  2419. index=categories,
  2420. columns=categories
  2421. )
  2422. # Drop 'O' row and column to get 3x3
  2423. lst_dom_bhvr_comb_3x3 = lst_dom_bhvr_comb_4x4.drop(index='O', columns='O')
  2424. # Normalize 3x3 matrix so that total sums to 100%
  2425. matrix_sum = lst_dom_bhvr_comb_3x3.values.sum()
  2426. if matrix_sum > 0:
  2427. lst_dom_bhvr_comb_3x3_normalized = lst_dom_bhvr_comb_3x3 / matrix_sum * 100
  2428. else:
  2429. lst_dom_bhvr_comb_3x3_normalized = lst_dom_bhvr_comb_3x3
  2430. # Save normalized 3x3 matrix
  2431. lst_dom_bhvr_comb_3x3_normalized.to_csv('data/Fig1L_dom_bhvr_comb.csv')
  2432. # %%
  2433. lst_dom_bhvr_comb_3x3_normalized = pd.read_csv('data/Fig1L_dom_bhvr_comb.csv', index_col=0)
  2434. plt.figure()
  2435. sns.heatmap(lst_dom_bhvr_comb_3x3_normalized, annot=True, fmt='.2f', cmap='viridis', vmin=0)
  2436. plt.gca().invert_yaxis()
  2437. plt.xlabel('Dom Behavior')
  2438. plt.ylabel('Sub Behavior')
  2439. plt.show()
  2440. # %%
  2441. from analyze_data_utils import analyze_independence
  2442. # Compute actual total_observations from original data
  2443. # Get number of dominant trials
  2444. num_dom_trials = len(lst_bhvr_frames_dict['p']['D'])
  2445. frames_per_trial = frame_end - frame_start # 300 - 150 = 150 frames
  2446. total_observations = num_dom_trials * frames_per_trial # total frames
  2447. print(f"Number of dominant trials: {num_dom_trials}")
  2448. print(f"Frames per trial: {frames_per_trial}")
  2449. print(f"Total observations (total_observations): {total_observations}")
  2450. # Convert percentages to counts
  2451. # lst_dom_bhvr_comb_3x3_normalized contains normalized percentages (sum=100%)
  2452. # Convert to actual frame counts
  2453. sub_dom_raw_behaviors = (lst_dom_bhvr_comb_3x3_normalized.values / 100 * total_observations).round().astype(int)
  2454. behavior_labels = lst_dom_bhvr_comb_3x3_normalized.index.tolist()
  2455. # Run full independence analysis
  2456. results_behaviors = analyze_independence(
  2457. sub_dom_raw=sub_dom_raw_behaviors,
  2458. labels=behavior_labels,
  2459. title="LST: Behavior Combinations Independence Test (Defense Time)",
  2460. data_type="time"
  2461. )
  2462. # %% [markdown]
  2463. # ### Fig1M. Defense time (%)
  2464. # %%
  2465. import numpy as np
  2466. import pandas as pd
  2467. from analyze_data_utils import calculate_total, filter_dict_data, filter_in_range, calculate_duration_time, dict_to_dataframe, merge_dicts
  2468. # Define framerate and time periods
  2469. framerate = 30
  2470. time_periods = {
  2471. 'baseline': {'range': [0, 4500], 'total_time': 150}, # 150s
  2472. '0-5s': {'range': [150, 300], 'total_time': 5}, # 5s
  2473. '5-20s': {'range': [300, 750], 'total_time': 15}, # 15s
  2474. '20-60s': {'range': [750, 1950], 'total_time': 40} # 40s
  2475. }
  2476. # Extract defense data (escape, freezing, rearing, tail-rattling) for baseline and looming periods
  2477. escape_data_bl = filter_dict_data(lst_baseline_bhvr_labels_dict, 'escape')
  2478. freezing_data_bl = filter_dict_data(lst_baseline_bhvr_labels_dict, 'freezing')
  2479. rearing_data_bl = filter_dict_data(lst_baseline_bhvr_labels_dict, 'rearing/up_stretch')
  2480. tailrattling_data_bl = filter_dict_data(lst_baseline_bhvr_labels_dict, 'tail_rattling')
  2481. escape_data = filter_dict_data(lst_bhvr_labels_dict, 'escape')
  2482. freezing_data = filter_dict_data(lst_bhvr_labels_dict, 'freezing')
  2483. rearing_data = filter_dict_data(lst_bhvr_labels_dict, 'rearing/up_stretch')
  2484. tailrattling_data = filter_dict_data(lst_bhvr_labels_dict, 'tail_rattling')
  2485. # Merge all defense behaviors
  2486. defense_data_bl = merge_dicts(escape_data_bl, freezing_data_bl)
  2487. defense_data_bl = merge_dicts(defense_data_bl, rearing_data_bl)
  2488. defense_data_bl = merge_dicts(defense_data_bl, tailrattling_data_bl)
  2489. defense_data = merge_dicts(escape_data, freezing_data)
  2490. defense_data = merge_dicts(defense_data, rearing_data)
  2491. defense_data = merge_dicts(defense_data, tailrattling_data)
  2492. # Store percentage data for all time periods
  2493. all_percentage_data = {}
  2494. # Process baseline period - use replace method to clip frames within range
  2495. defense_data_bl_filtered = filter_in_range(defense_data_bl, time_periods['baseline']['range'], method='replace')
  2496. duration_bl = calculate_duration_time(defense_data_bl_filtered, framerate=framerate)
  2497. duration_bl = calculate_total(duration_bl)
  2498. duration_bl_df = dict_to_dataframe(duration_bl, value_name='duration', nan2zero=True)
  2499. all_percentage_data['baseline'] = duration_bl_df
  2500. # Process 3 looming periods - use replace method to clip frames within range
  2501. for period_name in ['0-5s', '5-20s', '20-60s']:
  2502. period_range = time_periods[period_name]['range']
  2503. defense_period = filter_in_range(defense_data, period_range, method='replace')
  2504. duration_period = calculate_duration_time(defense_period, framerate=framerate)
  2505. duration_period = calculate_total(duration_period)
  2506. duration_period_df = dict_to_dataframe(duration_period, value_name='duration', nan2zero=True)
  2507. all_percentage_data[period_name] = duration_period_df
  2508. groups = ['s', 'p']
  2509. ranks = ['D', 'S']
  2510. # Store final data
  2511. final_data = {}
  2512. for period_name, df in all_percentage_data.items():
  2513. total_time = time_periods[period_name]['total_time']
  2514. row_data = []
  2515. # Extract data in the order: sD, sS, pD, pS
  2516. for group in groups:
  2517. for rank in ranks:
  2518. col_name = f'duration_{group}{rank}'
  2519. if col_name in df.columns:
  2520. # Calculate percentages and keep original order (sorted by B_id)
  2521. percentages = (df[col_name] / total_time) * 100
  2522. row_data.extend(percentages.tolist())
  2523. final_data[period_name] = row_data
  2524. # Create final DataFrame
  2525. final_df = pd.DataFrame(final_data).T
  2526. # Generate column names: sD_1, sD_2, ..., sS_1, sS_2, ..., pD_1, pD_2, ..., pS_1, pS_2, ...
  2527. column_names = []
  2528. first_df = list(all_percentage_data.values())[0]
  2529. for group in groups:
  2530. for rank in ranks:
  2531. col_name = f'duration_{group}{rank}'
  2532. if col_name in first_df.columns:
  2533. count = len(first_df)
  2534. column_names.extend([f'{group}{rank}_{i+1}' for i in range(count)])
  2535. final_df.columns = column_names
  2536. # Final result
  2537. defense_percentage_df = final_df
  2538. # Save result
  2539. defense_percentage_df.to_csv('data/Fig1M_defense_time_pct.csv')
  2540. # %%
  2541. import pandas as pd
  2542. import numpy as np
  2543. import matplotlib.pyplot as plt
  2544. df = pd.read_csv('data/Fig1M_defense_time_pct.csv')
  2545. time = df['Unnamed: 0'].values
  2546. groups = ['sD', 'sS', 'pD', 'pS']
  2547. fig, ax = plt.subplots()
  2548. style = {
  2549. 'sD': {'color': 'orange', 'linestyle': '--', 'marker': 'o', 'mfc': 'none'},
  2550. 'sS': {'color': 'blue', 'linestyle': '--', 'marker': 'o', 'mfc': 'none'},
  2551. 'pD': {'color': 'orange', 'linestyle': '-', 'marker': 'o', 'mfc': 'orange'},
  2552. 'pS': {'color': 'blue', 'linestyle': '-', 'marker': 'o', 'mfc': 'blue'}
  2553. }
  2554. for g in groups:
  2555. cols = [c for c in df.columns if c.startswith(g + '_')]
  2556. means = df[cols].mean(axis=1)
  2557. stds = df[cols].sem(axis=1)
  2558. ax.errorbar(
  2559. time,
  2560. means,
  2561. yerr=stds,
  2562. color=style[g]['color'],
  2563. linestyle=style[g]['linestyle'],
  2564. marker=style[g]['marker'],
  2565. markerfacecolor=style[g]['mfc'],
  2566. capsize=4,
  2567. linewidth=2,
  2568. label=g
  2569. )
  2570. ax.set_xticks(range(len(time)))
  2571. ax.set_xticklabels(time, rotation=45, ha='right')
  2572. ax.set_ylabel('Percentage')
  2573. ax.set_ylim(0, 60)
  2574. ax.legend()
  2575. plt.tight_layout()
  2576. plt.show()
  2577. # %% [markdown]
  2578. # ### Fig1N. Peak escape speed (cm/s)
  2579. # %%
  2580. import numpy as np
  2581. from analyze_data_utils import filter_dict_data, filter_in_range, calculate_latency_time, dict_to_dataframe
  2582. escape_data = filter_dict_data(lst_bhvr_labels_dict, 'escape')
  2583. escape_data = filter_in_range(escape_data, [150,1950], method='delete')
  2584. escape_latency = calculate_latency_time(escape_data)
  2585. escape_latency_df = dict_to_dataframe(escape_latency, value_name='escape_latency')
  2586. escape_latency_df.to_csv('data/lst_1st_escape_latency.csv', index=False)
  2587. vel_escape_data = {'p':{'D': {}, 'S': {}}, 's': {'D': {}, 'S': {}}}
  2588. vel_escape_latency = {'p':{'D': {}, 'S': {}}, 's': {'D': {}, 'S': {}}}
  2589. for g in ['s', 'p']:
  2590. for r in ['D', 'S']:
  2591. for t_id in escape_data[g][r].keys():
  2592. if t_id.split('B')[1].split('M')[0] in ['1', '3', '7', '10', '11', '12', '13', '14', '15', '16', '17', '20']:
  2593. m = 1 if t_id.split('M')[1].split('T')[0] == 'S' else 2
  2594. else:
  2595. m = 1 if t_id.split('M')[1].split('T')[0] == 'D' else 2
  2596. m_id = t_id.split('M')[0] + 'M' + str(m) + '(' + t_id.split('M')[1].split('T')[0] + ')'
  2597. if g == 's':
  2598. velocity_data = lsts_data[m_id]['velocity']['waist']
  2599. looming_start_frame = looming_start_frames[m_id][int(t_id.split('T')[1])-1]
  2600. elif g == 'p':
  2601. velocity_data = lstp_data[m_id]['velocity']['waist']
  2602. mp_id = mp_ids[int(m_id.split('B')[1].split('M')[0])-1]
  2603. looming_start_frame = looming_start_frames[mp_id][int(t_id.split('T')[1])-1]
  2604. escape_list = escape_data[g][r][t_id]
  2605. for escape in escape_list:
  2606. if np.any(np.isnan(escape)):
  2607. continue
  2608. else:
  2609. start, end = escape
  2610. escape_start_frame = looming_start_frame+start-150
  2611. escape_end_frame = looming_start_frame+end-150
  2612. vel_nest = velocity_data[escape_start_frame:escape_end_frame].max() / lst_pixpercm * framerate
  2613. if t_id not in vel_escape_data[g][r].keys():
  2614. vel_escape_data[g][r][t_id] = []
  2615. vel_escape_data[g][r][t_id].append(vel_nest)
  2616. max_velocity_index = np.argmax(velocity_data[escape_start_frame:escape_end_frame])
  2617. max_velocity_frame = escape_start_frame + max_velocity_index
  2618. latency_to_max_velocity = (max_velocity_frame - escape_start_frame) / framerate
  2619. if t_id in ['B10MDT3', 'B10MDT4', 'B10MDT5', 'B10MDT9']:
  2620. # Get the velocity segment for these specific trials
  2621. velocity_segment = velocity_data[escape_start_frame:escape_end_frame].values
  2622. velocity_segment_cm_s = velocity_segment / lst_pixpercm * framerate
  2623. if t_id not in vel_escape_latency[g][r].keys():
  2624. vel_escape_latency[g][r][t_id] = []
  2625. vel_escape_latency[g][r][t_id].append(latency_to_max_velocity)
  2626. vel_escape_df = dict_to_dataframe(vel_escape_data, value_name='escape_speed')
  2627. vel_escape_df.to_csv('data/Fig1N_max_escape_speed.csv', index=False)
  2628. # %%
  2629. import pandas as pd
  2630. import numpy as np
  2631. import matplotlib.pyplot as plt
  2632. df = pd.read_csv('data/Fig1N_max_escape_speed.csv')
  2633. groups = ['escape_speed_sD', 'escape_speed_pD', 'escape_speed_sS', 'escape_speed_pS']
  2634. fig, ax = plt.subplots()
  2635. x_pos = {}
  2636. xticks = []
  2637. xticklabels = []
  2638. x = 0
  2639. colors = {'sD': 'orange', 'sS': 'blue', 'pD': 'orange', 'pS': 'blue'}
  2640. # 1. x
  2641. for g in groups:
  2642. cond = g.split('_')[-1]
  2643. if cond in ['sD', 'sS']:
  2644. x_pos[f'{g}_mean'] = x
  2645. x_pos[f'{g}_pts'] = x + 0.8
  2646. xticks += [x, x + 0.8]
  2647. xticklabels += [f'{g}_mean', f'{g}_pts']
  2648. else:
  2649. x_pos[f'{g}_pts'] = x
  2650. x_pos[f'{g}_mean'] = x + 0.8
  2651. xticks += [x, x + 0.8]
  2652. xticklabels += [f'{g}_pts', f'{g}_mean']
  2653. x += 2.2
  2654. # 2. mean ± SD
  2655. for g in groups:
  2656. cond = g.split('_')[-1]
  2657. ax.errorbar(
  2658. x_pos[f'{g}_mean'],
  2659. df[g].mean(),
  2660. yerr=df[g].std(),
  2661. fmt='o',
  2662. color=colors[cond],
  2663. capsize=4,
  2664. markersize=9,
  2665. elinewidth=1.5,
  2666. capthick=1.5,
  2667. zorder=3
  2668. )
  2669. # 3. raw + correct pairing
  2670. for i in range(len(df)):
  2671. sD_x = x_pos['escape_speed_sD_pts']
  2672. sS_x = x_pos['escape_speed_sS_pts']
  2673. pD_x = x_pos['escape_speed_pD_pts']
  2674. pS_x = x_pos['escape_speed_pS_pts']
  2675. sD_y = df.loc[i, 'escape_speed_sD']
  2676. sS_y = df.loc[i, 'escape_speed_sS']
  2677. pD_y = df.loc[i, 'escape_speed_pD']
  2678. pS_y = df.loc[i, 'escape_speed_pS']
  2679. ax.scatter(sD_x, sD_y, color=colors['sD'])
  2680. ax.scatter(sS_x, sS_y, color=colors['sS'])
  2681. ax.scatter(pD_x, pD_y, color=colors['pD'])
  2682. ax.scatter(pS_x, pS_y, color=colors['pS'])
  2683. ax.plot([sD_x, pD_x], [sD_y, pD_y], color='gray', linewidth=1)
  2684. ax.plot([sS_x, pS_x], [sS_y, pS_y], color='gray', linewidth=1)
  2685. # 4. axis
  2686. ax.set_xticks(xticks)
  2687. ax.set_xticklabels(xticklabels, rotation=45, ha='right')
  2688. ax.set_ylabel('Peak escape speed (cm/s)')
  2689. ax.set_ylim(0, 125)
  2690. plt.tight_layout()
  2691. plt.show()
  2692. # %% [markdown]
  2693. # ### Fig1O. First freezing duration (s)
  2694. # %%
  2695. from analyze_data_utils import filter_dict_data, filter_duration, calculate_duration_time
  2696. def get_behavior_frame_when_main_behavior(lst_deci_dict, lst_bhvr_labels_dict, main_behavior, sub_behavior=None):
  2697. target_behavior_data = {'p':{'D': {}, 'S': {}}, 's': {'D': {}, 'S': {}}}
  2698. for g in lst_bhvr_labels_dict.keys():
  2699. for m in lst_bhvr_labels_dict[g].keys():
  2700. for t in lst_bhvr_labels_dict[g][m].keys():
  2701. if lst_deci_dict[g][m][t] in main_behavior:
  2702. if sub_behavior:
  2703. sub_behavior_data = lst_bhvr_labels_dict[g][m][t][sub_behavior]
  2704. else:
  2705. sub_behavior_data = lst_bhvr_labels_dict[g][m][t]
  2706. target_behavior_data[g][m][t] = sub_behavior_data
  2707. return target_behavior_data
  2708. def calculate_bhvr1_duration_before_bhvr2(bhvr1, bhvr2, first=False):
  2709. duration_data = {}
  2710. for key, values in bhvr1.items():
  2711. if isinstance(values, dict):
  2712. duration_data[key] = calculate_bhvr1_duration_before_bhvr2(bhvr1[key], bhvr2[key], first=first)
  2713. else:
  2714. turples = []
  2715. if not np.isnan(values).any():
  2716. for value in values:
  2717. start, end = value
  2718. if np.isnan(bhvr2[key]).any():
  2719. turples.append((start, end))
  2720. else:
  2721. if end <= bhvr2[key][0][0]+1: # 使用第一次escape的开始帧
  2722. turples.append((start, end))
  2723. # 如果需要第一次且turples不为空,只保留第一个
  2724. if first and len(turples) > 0:
  2725. turples = [turples[0]]
  2726. durations = []
  2727. if np.isnan(turples).any():
  2728. duration = np.nan
  2729. durations.append(duration)
  2730. duration_data[key] = durations
  2731. else:
  2732. for turple in turples:
  2733. start, end = turple
  2734. duration = (end - start + 1) / framerate # +1 for inclusive counting
  2735. durations.append(duration)
  2736. duration_data[key] = durations
  2737. return duration_data
  2738. freezing_data = get_behavior_frame_when_main_behavior(lst_deci_dict, lst_bhvr_labels_dict, ['F', 'F+E'], 'freezing')
  2739. # freezing_data = filter_dict_data(behavior_frames_dict, 'freezing')
  2740. escape_data = filter_dict_data(lst_bhvr_labels_dict, 'escape')
  2741. freezing_data = filter_duration(freezing_data)
  2742. freezing_data = filter_in_range(freezing_data, [150, 300], method='replace')
  2743. first_freezing_duration = calculate_bhvr1_duration_before_bhvr2(freezing_data, escape_data, first=True)
  2744. first_freezing_duration_df = dict_to_dataframe(first_freezing_duration, value_name='first_freezing_duration')
  2745. first_freezing_duration_df.to_csv('data/Fig1O_first_freezing_duration.csv', index=False)
  2746. # %%
  2747. import pandas as pd
  2748. import numpy as np
  2749. import matplotlib.pyplot as plt
  2750. df = pd.read_csv('data/Fig1O_first_freezing_duration.csv')
  2751. print(df.columns.tolist())
  2752. groups = ['first_freezing_duration_sD', 'first_freezing_duration_pD', 'first_freezing_duration_sS', 'first_freezing_duration_pS']
  2753. fig, ax = plt.subplots()
  2754. x_pos = {}
  2755. xticks = []
  2756. xticklabels = []
  2757. x = 0
  2758. colors = {'sD': 'orange', 'sS': 'blue', 'pD': 'orange', 'pS': 'blue'}
  2759. # 1. x
  2760. for g in groups:
  2761. cond = g.split('_')[-1]
  2762. if cond in ['sD', 'sS']:
  2763. x_pos[f'{g}_mean'] = x
  2764. x_pos[f'{g}_pts'] = x + 0.8
  2765. xticks += [x, x + 0.8]
  2766. xticklabels += [f'{g}_mean', f'{g}_pts']
  2767. else:
  2768. x_pos[f'{g}_pts'] = x
  2769. x_pos[f'{g}_mean'] = x + 0.8
  2770. xticks += [x, x + 0.8]
  2771. xticklabels += [f'{g}_pts', f'{g}_mean']
  2772. x += 2.2
  2773. # 2. mean ± SD
  2774. for g in groups:
  2775. cond = g.split('_')[-1]
  2776. ax.errorbar(
  2777. x_pos[f'{g}_mean'],
  2778. df[g].mean(),
  2779. yerr=df[g].std(),
  2780. fmt='o',
  2781. color=colors[cond],
  2782. capsize=4,
  2783. markersize=9,
  2784. elinewidth=1.5,
  2785. capthick=1.5,
  2786. zorder=3
  2787. )
  2788. # 3. raw + correct pairing
  2789. for i in range(len(df)):
  2790. sD_x = x_pos['first_freezing_duration_sD_pts']
  2791. sS_x = x_pos['first_freezing_duration_sS_pts']
  2792. pD_x = x_pos['first_freezing_duration_pD_pts']
  2793. pS_x = x_pos['first_freezing_duration_pS_pts']
  2794. sD_y = df.loc[i, 'first_freezing_duration_sD']
  2795. sS_y = df.loc[i, 'first_freezing_duration_sS']
  2796. pD_y = df.loc[i, 'first_freezing_duration_pD']
  2797. pS_y = df.loc[i, 'first_freezing_duration_pS']
  2798. ax.scatter(sD_x, sD_y, color=colors['sD'])
  2799. ax.scatter(sS_x, sS_y, color=colors['sS'])
  2800. ax.scatter(pD_x, pD_y, color=colors['pD'])
  2801. ax.scatter(pS_x, pS_y, color=colors['pS'])
  2802. ax.plot([sD_x, pD_x], [sD_y, pD_y], color='gray', linewidth=1)
  2803. ax.plot([sS_x, pS_x], [sS_y, pS_y], color='gray', linewidth=1)
  2804. # 4. axis
  2805. ax.set_xticks(xticks)
  2806. ax.set_xticklabels(xticklabels, rotation=45, ha='right')
  2807. ax.set_ylabel('1st freezing duration (s)')
  2808. ax.set_ylim(0, 4)
  2809. plt.tight_layout()
  2810. plt.show()
  2811. # %% [markdown]
  2812. # ### Fig1P. Grooming time (%)
  2813. # %%
  2814. import numpy as np
  2815. import pandas as pd
  2816. from analyze_data_utils import calculate_total, filter_dict_data, filter_in_range, calculate_duration_time, dict_to_dataframe
  2817. # Define framerate and time periods
  2818. framerate = 30
  2819. time_periods = {
  2820. 'baseline': {'range': [0, 4500], 'total_time': 150}, # 150s
  2821. '0-5s': {'range': [150, 300], 'total_time': 5}, # 5s
  2822. '5-20s': {'range': [300, 750], 'total_time': 15}, # 15s
  2823. '20-60s': {'range': [750, 1950], 'total_time': 40} # 40s
  2824. }
  2825. # Extract grooming data for baseline and looming periods
  2826. grooming_data_bl = filter_dict_data(lst_baseline_bhvr_labels_dict, 'grooming')
  2827. grooming_data = filter_dict_data(lst_bhvr_labels_dict, 'grooming')
  2828. # Store percentage data for all time periods
  2829. all_percentage_data = {}
  2830. # Process baseline period - use replace method to clip frames within range
  2831. grooming_data_bl_filtered = filter_in_range(grooming_data_bl, time_periods['baseline']['range'], method='replace')
  2832. duration_bl = calculate_duration_time(grooming_data_bl_filtered, framerate=framerate)
  2833. duration_bl = calculate_total(duration_bl)
  2834. duration_bl_df = dict_to_dataframe(duration_bl, value_name='duration', nan2zero=True)
  2835. all_percentage_data['baseline'] = duration_bl_df
  2836. # Process 3 looming periods - use replace method to clip frames within range
  2837. for period_name in ['0-5s', '5-20s', '20-60s']:
  2838. period_range = time_periods[period_name]['range']
  2839. grooming_period = filter_in_range(grooming_data, period_range, method='replace')
  2840. duration_period = calculate_duration_time(grooming_period, framerate=framerate)
  2841. duration_period = calculate_total(duration_period)
  2842. duration_period_df = dict_to_dataframe(duration_period, value_name='duration', nan2zero=True)
  2843. all_percentage_data[period_name] = duration_period_df
  2844. groups = ['s', 'p']
  2845. ranks = ['D', 'S']
  2846. # Store final data
  2847. final_data = {}
  2848. for period_name, df in all_percentage_data.items():
  2849. total_time = time_periods[period_name]['total_time']
  2850. row_data = []
  2851. # Extract data in the order: sD, sS, pD, pS
  2852. for group in groups:
  2853. for rank in ranks:
  2854. col_name = f'duration_{group}{rank}'
  2855. if col_name in df.columns:
  2856. # Calculate percentages and keep original order (sorted by B_id)
  2857. percentages = (df[col_name] / total_time) * 100
  2858. row_data.extend(percentages.tolist())
  2859. final_data[period_name] = row_data
  2860. # Create final DataFrame
  2861. final_df = pd.DataFrame(final_data).T
  2862. # Generate column names: sD_1, sD_2, ..., sS_1, sS_2, ..., pD_1, pD_2, ..., pS_1, pS_2, ...
  2863. column_names = []
  2864. first_df = list(all_percentage_data.values())[0]
  2865. for group in groups:
  2866. for rank in ranks:
  2867. col_name = f'duration_{group}{rank}'
  2868. if col_name in first_df.columns:
  2869. count = len(first_df)
  2870. column_names.extend([f'{group}{rank}_{i+1}' for i in range(count)])
  2871. final_df.columns = column_names
  2872. # Final result
  2873. grooming_percentage_df = final_df
  2874. # Save result
  2875. grooming_percentage_df.to_csv('data/Fig1P_grooming_time_pct.csv')
  2876. # %%
  2877. import pandas as pd
  2878. import numpy as np
  2879. import matplotlib.pyplot as plt
  2880. df = pd.read_csv('data/Fig1P_grooming_time_pct.csv')
  2881. time = df['Unnamed: 0'].values
  2882. groups = ['sD', 'sS', 'pD', 'pS']
  2883. fig, ax = plt.subplots()
  2884. style = {
  2885. 'sD': {'color': 'orange', 'linestyle': '--', 'marker': 'o', 'mfc': 'none'},
  2886. 'sS': {'color': 'blue', 'linestyle': '--', 'marker': 'o', 'mfc': 'none'},
  2887. 'pD': {'color': 'orange', 'linestyle': '-', 'marker': 'o', 'mfc': 'orange'},
  2888. 'pS': {'color': 'blue', 'linestyle': '-', 'marker': 'o', 'mfc': 'blue'}
  2889. }
  2890. for g in groups:
  2891. cols = [c for c in df.columns if c.startswith(g + '_')]
  2892. means = df[cols].mean(axis=1)
  2893. stds = df[cols].sem(axis=1)
  2894. ax.errorbar(
  2895. time,
  2896. means,
  2897. yerr=stds,
  2898. color=style[g]['color'],
  2899. linestyle=style[g]['linestyle'],
  2900. marker=style[g]['marker'],
  2901. markerfacecolor=style[g]['mfc'],
  2902. capsize=4,
  2903. linewidth=2,
  2904. label=g
  2905. )
  2906. ax.set_xticks(range(len(time)))
  2907. ax.set_xticklabels(time, rotation=45, ha='right')
  2908. ax.set_ylabel('Percentage')
  2909. ax.set_ylim(0, 20)
  2910. ax.legend()
  2911. plt.tight_layout()
  2912. plt.show()
  2913. # %% [markdown]
  2914. # ### Fig1Q. Rearing & up-stretch time (%)
  2915. # %%
  2916. import numpy as np
  2917. import pandas as pd
  2918. from analyze_data_utils import calculate_total, filter_dict_data, filter_in_range, calculate_duration_time, dict_to_dataframe
  2919. # Define framerate and time periods
  2920. framerate = 30
  2921. time_periods = {
  2922. 'baseline': {'range': [0, 4500], 'total_time': 150}, # 150s
  2923. '0-5s': {'range': [150, 300], 'total_time': 5}, # 5s
  2924. '5-20s': {'range': [300, 750], 'total_time': 15}, # 15s
  2925. '20-60s': {'range': [750, 1950], 'total_time': 40} # 40s
  2926. }
  2927. # Extract rearing/up_stretch data for baseline and looming periods
  2928. rearing_data_bl = filter_dict_data(lst_baseline_bhvr_labels_dict, 'rearing/up_stretch')
  2929. rearing_data = filter_dict_data(lst_bhvr_labels_dict, 'rearing/up_stretch')
  2930. # Store percentage data for all time periods
  2931. all_percentage_data = {}
  2932. # Process baseline period - use replace method to clip frames within range
  2933. rearing_data_bl_filtered = filter_in_range(rearing_data_bl, time_periods['baseline']['range'], method='replace')
  2934. duration_bl = calculate_duration_time(rearing_data_bl_filtered, framerate=framerate)
  2935. duration_bl = calculate_total(duration_bl)
  2936. duration_bl_df = dict_to_dataframe(duration_bl, value_name='duration', nan2zero=True)
  2937. all_percentage_data['baseline'] = duration_bl_df
  2938. # Process 3 looming periods - use replace method to clip frames within range
  2939. for period_name in ['0-5s', '5-20s', '20-60s']:
  2940. period_range = time_periods[period_name]['range']
  2941. rearing_period = filter_in_range(rearing_data, period_range, method='replace')
  2942. duration_period = calculate_duration_time(rearing_period, framerate=framerate)
  2943. duration_period = calculate_total(duration_period)
  2944. duration_period_df = dict_to_dataframe(duration_period, value_name='duration', nan2zero=True)
  2945. all_percentage_data[period_name] = duration_period_df
  2946. groups = ['s', 'p']
  2947. ranks = ['D', 'S']
  2948. # Store final data
  2949. final_data = {}
  2950. for period_name, df in all_percentage_data.items():
  2951. total_time = time_periods[period_name]['total_time']
  2952. row_data = []
  2953. # Extract data in the order: sD, sS, pD, pS
  2954. for group in groups:
  2955. for rank in ranks:
  2956. col_name = f'duration_{group}{rank}'
  2957. if col_name in df.columns:
  2958. # Calculate percentages and keep original order (sorted by B_id)
  2959. percentages = (df[col_name] / total_time) * 100
  2960. row_data.extend(percentages.tolist())
  2961. final_data[period_name] = row_data
  2962. # Create final DataFrame
  2963. final_df = pd.DataFrame(final_data).T
  2964. # Generate column names: sD_1, sD_2, ..., sS_1, sS_2, ..., pD_1, pD_2, ..., pS_1, pS_2, ...
  2965. column_names = []
  2966. first_df = list(all_percentage_data.values())[0]
  2967. for group in groups:
  2968. for rank in ranks:
  2969. col_name = f'duration_{group}{rank}'
  2970. if col_name in first_df.columns:
  2971. count = len(first_df)
  2972. column_names.extend([f'{group}{rank}_{i+1}' for i in range(count)])
  2973. final_df.columns = column_names
  2974. # Save result
  2975. rearing_percentage_df = final_df
  2976. rearing_percentage_df.to_csv('data/Fig1Q_rearing_up_stretch_time_pct.csv')
  2977. # %%
  2978. import pandas as pd
  2979. import numpy as np
  2980. import matplotlib.pyplot as plt
  2981. df = pd.read_csv('data/Fig1Q_rearing_up_stretch_time_pct.csv')
  2982. time = df['Unnamed: 0'].values
  2983. groups = ['sD', 'sS', 'pD', 'pS']
  2984. fig, ax = plt.subplots()
  2985. style = {
  2986. 'sD': {'color': 'orange', 'linestyle': '--', 'marker': 'o', 'mfc': 'none'},
  2987. 'sS': {'color': 'blue', 'linestyle': '--', 'marker': 'o', 'mfc': 'none'},
  2988. 'pD': {'color': 'orange', 'linestyle': '-', 'marker': 'o', 'mfc': 'orange'},
  2989. 'pS': {'color': 'blue', 'linestyle': '-', 'marker': 'o', 'mfc': 'blue'}
  2990. }
  2991. for g in groups:
  2992. cols = [c for c in df.columns if c.startswith(g + '_')]
  2993. means = df[cols].mean(axis=1)
  2994. stds = df[cols].sem(axis=1)
  2995. ax.errorbar(
  2996. time,
  2997. means,
  2998. yerr=stds,
  2999. color=style[g]['color'],
  3000. linestyle=style[g]['linestyle'],
  3001. marker=style[g]['marker'],
  3002. markerfacecolor=style[g]['mfc'],
  3003. capsize=4,
  3004. linewidth=2,
  3005. label=g
  3006. )
  3007. ax.set_xticks(range(len(time)))
  3008. ax.set_xticklabels(time, rotation=45, ha='right')
  3009. ax.set_ylabel('Percentage')
  3010. ax.set_ylim(0, 20)
  3011. ax.legend()
  3012. plt.tight_layout()
  3013. plt.show()
  3014. # %% [markdown]
  3015. # ### FigS1A. Time in zone when different edge_width
  3016. # %%
  3017. import matplotlib.pyplot as plt
  3018. def count_sequence(series):
  3019. count = 0
  3020. target_structures = [
  3021. [1, 1, 1, 1, 1, 2, 2, 2, 2, 2],
  3022. [2, 2, 2, 2, 2, 1, 1, 1, 1, 1],
  3023. ]
  3024. current_structure = []
  3025. approach_index = []
  3026. for i, value in enumerate(series):
  3027. current_structure.append(value)
  3028. if len(current_structure) > 10:
  3029. current_structure = current_structure[1:]
  3030. if current_structure in target_structures:
  3031. count += 1
  3032. if value == 1:
  3033. approach_index.append(i)
  3034. return count, approach_index
  3035. center_pcts_list = []
  3036. edge_pcts_list = []
  3037. nest_pcts_list = []
  3038. trans_bouts_list = []
  3039. for i in range(0, 26):
  3040. edge_width = 43.52*i/2
  3041. center_pcts = []
  3042. edge_pcts = []
  3043. nest_pcts = []
  3044. trans_bouts = []
  3045. for m_id in ms_ids:
  3046. coord_data = {
  3047. 'x': pd.to_numeric(lsts_data[m_id]['coordinate']['centroid']['x'][0:18000], errors='coerce'),
  3048. 'y': pd.to_numeric(lsts_data[m_id]['coordinate']['centroid']['y'][0:18000], errors='coerce')}
  3049. coord_df = pd.DataFrame(coord_data).dropna()
  3050. local_data = get_lst_location_value(coord_data, edge_width)
  3051. center_pcts.append(local_data.value_counts().get(2, 0) / len(local_data) * 100)
  3052. edge_pcts.append(local_data.value_counts().get(1, 0) / len(local_data) * 100)
  3053. nest_pcts.append(local_data.value_counts().get(0, 0) / len(local_data) * 100)
  3054. trans_bout, trans_idx = count_sequence(local_data)
  3055. trans_bouts.append(trans_bout)
  3056. center_pcts_list.append(np.mean(center_pcts))
  3057. edge_pcts_list.append(np.mean(edge_pcts))
  3058. nest_pcts_list.append(np.mean(nest_pcts))
  3059. trans_bouts_list.append(np.mean(trans_bouts))
  3060. # Calculate edge widths in centimeters for each percentage data point
  3061. edge_widths_cm = [43.52 * i / 2 for i in range(len(edge_pcts_list))]
  3062. df_edge_analysis = pd.DataFrame({
  3063. 'edge_width_cm': edge_widths_cm,
  3064. 'edge_pct': edge_pcts_list,
  3065. 'trans_bouts': trans_bouts_list
  3066. })
  3067. df_edge_analysis.to_csv('data/FigS1A_edge_zone_width.csv', index=False)
  3068. # %%
  3069. from kneed import KneeLocator
  3070. df = pd.read_csv('data/FigS1A_edge_zone_width.csv')
  3071. edge_pcts_list = df['edge_pct'].values
  3072. trans_bouts_list = df['trans_bouts'].values
  3073. # Find the knee point using the Kneedle algorithm
  3074. edge_widths_kneed = np.arange(len(edge_pcts_list))
  3075. edge_pcts_kneed = np.array(edge_pcts_list)
  3076. kl = KneeLocator(
  3077. edge_widths_kneed, edge_pcts_kneed,
  3078. curve='concave', direction='increasing',
  3079. online=True
  3080. )
  3081. x0, x1 = edge_widths_kneed.min(), edge_widths_kneed.max()
  3082. x_diff_mapped = kl.x_difference * (x1 - x0) + x0
  3083. y_diff = kl.y_difference
  3084. fig, ax1 = plt.subplots(figsize=(12, 6), dpi=300)
  3085. ax1.plot(edge_pcts_list, color='red', label='Time in Edge Zone')
  3086. # ax1.plot(np.gradient(edge_pcts_list), color='orange', label='ΔTime in Edge Zone')
  3087. # ax1.plot(np.gradient(np.gradient(edge_pcts_list)), color='yellow', label='Δ²Time in Edge Zone')
  3088. ax1.axvline(x=kl.knee, color='green', linestyle='--', linewidth=1.5, label='Knee')
  3089. ax1.plot(x_diff_mapped, y_diff * np.max(edge_pcts_list),
  3090. color='limegreen', linestyle='-.',
  3091. label='Kneedle Difference')
  3092. ax1.set_ylabel('Time in zone (%)')
  3093. ax1.set_xlabel('Edge Width (cm)')
  3094. ax1.grid(True)
  3095. ax1.legend(loc='upper left')
  3096. ax2 = ax1.twinx()
  3097. ax2.plot(trans_bouts_list, color='darkgray', label='Edge <=> Center Transition Bouts')
  3098. ax2.set_ylabel('Transition Bouts')
  3099. ax2.legend(loc='upper right')
  3100. plt.tight_layout()
  3101. plt.show()
  3102. # %% [markdown]
  3103. # ### FigS1C. baseline ethogram
  3104. # %%
  3105. import pandas as pd
  3106. import matplotlib.pyplot as plt
  3107. import matplotlib.colors as mcolors
  3108. import colorsys
  3109. from matplotlib.patches import FancyArrow
  3110. from analyze_data_utils import filter_in_range
  3111. def adjust_saturation(color, saturation):
  3112. rgb = mcolors.to_rgb(color)
  3113. h, l, s = colorsys.rgb_to_hls(*rgb)
  3114. new_s = saturation * s
  3115. new_rgb = colorsys.hls_to_rgb(h, l, new_s)
  3116. return new_rgb
  3117. with open('data/lst_bl_etho_dict.pkl', 'rb') as f:
  3118. lst_bl_etho_dict = pickle.load(f)
  3119. lst_framerange = [0, 4500]
  3120. saturation = 0.9
  3121. framerate=30
  3122. stim_frame = int(0.725 * framerate)
  3123. behavior_frames_dict = lst_etho_dict
  3124. behavior_frames_dict = dict(sorted(behavior_frames_dict.items()))
  3125. lst_behavior_frames_dict = filter_in_range(lst_etho_dict, lst_framerange, method='replace')
  3126. lst_behavior_properties = {
  3127. 'approach_partner': ('#143FCA', 4),
  3128. 'follow_partner': ('#137CAB', 4),
  3129. 'groom_partner': ('#0990FF', 4),
  3130. 'sniff_partner': ('#09FFFF', 4),
  3131. 'tailrattling': ('#FF09FF', 5),
  3132. 'huddling': ('#6D57F3', 4),
  3133. 'jumping': ('#FF1717', 5),
  3134. 'escape': ('#FF1717', 4),
  3135. 'freezing': ('#FF75EF', 4),
  3136. 'dwelling': ('#FF75EF', 4),
  3137. 'grooming': ('#28AE61', 4),
  3138. 'real_rearing': ('#C409FF', 3),
  3139. 'stretching': ('#C409FF', 3),
  3140. 'nest': ('gray', 2),
  3141. 'others': ('white', 2),
  3142. 'rearing': ('white', 1),
  3143. 'sniffing': ('#0C8140', 1),
  3144. 'looming': ('white', 1),
  3145. 'reaction': ('white', 1),
  3146. 'climbing': ('gray', 1),
  3147. 'in_proximity': ('white', 1)
  3148. }
  3149. fig, axes = plt.subplots(1, 2, figsize=(10, 2), dpi=300)
  3150. plt.subplots_adjust(wspace=0)
  3151. for i, group in enumerate(lst_bl_etho_dict.keys()):
  3152. ax = axes[i]
  3153. ax.set_yticks(range(len(lst_bl_etho_dict[group].keys())))
  3154. for idx, t_id in enumerate(lst_bl_etho_dict[group].keys()):
  3155. behaviors = lst_bl_etho_dict[group][t_id]
  3156. for b, frame_ranges in behaviors.items():
  3157. for frame_range in frame_ranges:
  3158. if np.isnan(frame_ranges).all():
  3159. continue
  3160. start_frame, end_frame = frame_range
  3161. color, zorder = lst_behavior_properties.get(b, ('black', 0))
  3162. color = adjust_saturation(color, saturation)
  3163. rect = plt.Rectangle((start_frame, idx - 0.4), end_frame - start_frame, 0.8, facecolor=color, edgecolor='none', zorder=zorder)
  3164. ax.add_patch(rect)
  3165. if idx % 2 == 0:
  3166. ax.axhline(idx - 0.5, color='black', linestyle='-', linewidth=1, zorder=9)
  3167. ytick_color = 'red'
  3168. else:
  3169. ax.axhline(idx - 0.5, color='black', linestyle='-', linewidth=0.1, zorder=9)
  3170. ytick_color = 'green'
  3171. ax.get_yticklabels()[idx].set_color(ytick_color)
  3172. if i == 0:
  3173. ax.set_yticklabels(['D', 'S', 'D', 'S', 'D', 'S'])
  3174. ax.set_ylabel('Looming')
  3175. else:
  3176. ax.set_yticks([])
  3177. ax.set_xticks([])
  3178. lst_blank_width = (lst_framerange[1] - lst_framerange[0]) / 20
  3179. per_rect = plt.Rectangle((lst_framerange[0] - lst_blank_width, -0.5), lst_blank_width, 6,
  3180. facecolor='white', edgecolor='none', zorder=8)
  3181. ax.add_patch(per_rect)
  3182. post_rect = plt.Rectangle((lst_framerange[1], -0.5), lst_blank_width, 6,
  3183. facecolor='white', edgecolor='none', zorder=8)
  3184. ax.add_patch(post_rect)
  3185. ax.set_xlim(lst_framerange[0] - lst_blank_width, lst_framerange[1] + lst_blank_width)
  3186. ax.set_ylim(-0.5, 5.5)
  3187. ax.invert_yaxis()
  3188. ax.set_title(f'{group}')
  3189. if i == 1:
  3190. line_x1 = lst_framerange[1] - 30*30
  3191. line_x2 = lst_framerange[1]
  3192. ax.text((line_x1 + line_x2) / 2, -0.8, '30 sec', fontsize=8, color='black', ha='center')
  3193. lst_line = plt.Line2D([line_x1, line_x2], [-0.65, -0.65], color='black', linewidth=1)
  3194. lst_line.set_clip_on(False)
  3195. ax.add_artist(lst_line)
  3196. lst_legend_ = {
  3197. 'freezing': ('#FF75EF', 4),
  3198. 'tail rattling': ('#FF09FF', 5),
  3199. 'stretching / rearing': ('#C409FF', 3),
  3200. 'huddling': ('#6D57F3', 1),
  3201. 'approaching P (partner)': ('#143FCA', 4),
  3202. 'following P': ('#137CAB', 4),
  3203. 'grooming P': ('#0990FF', 4),
  3204. 'sniffing P': ('#09FFFF', 4),
  3205. 'grooming': ('#28AE61', 4),
  3206. 'sniffing': ('#0C8140', 1),
  3207. 'in nest': ('gray', 2),
  3208. 'other behaviors': ('white', 1)}
  3209. lst_legend = []
  3210. for label, (color, zorder) in lst_legend_.items():
  3211. if label == 'other behaviors':
  3212. rect = plt.Rectangle((0, 0), 1, 1, facecolor=adjust_saturation(color, saturation),
  3213. edgecolor='black', linewidth=0.5, label=label)
  3214. else:
  3215. rect = plt.Rectangle((0, 0), 1, 1, facecolor=adjust_saturation(color, saturation), label=label)
  3216. lst_legend.append(rect)
  3217. axes[1].legend(handles=lst_legend, loc='upper center', bbox_to_anchor=(0, -0.2), fontsize=7, ncol=7)
  3218. plt.savefig('fig/fig_1/lst_baseline_behavior_ethogram.eps', format="eps", dpi=300, bbox_inches="tight")
  3219. plt.show()
  3220. # %% [markdown]
  3221. # ### FigS1D. Co-trigger percentage
  3222. # %%
  3223. import pandas as pd
  3224. lst_decision_file = os.path.abspath(os.path.join('.', 'lst', 'looming_behavior_sb_1looming_new.csv'))
  3225. trial_data = pd.read_csv(lst_decision_file, header=None, nrows=340)
  3226. # Groups where M1 is MS and M2 is MD
  3227. md_is_m2_groups = ['1', '3', '7', '10', '11', '12', '13', '14', '15', '16', '17', '20']
  3228. md_trigger_count = 0 # only MD triggered
  3229. ms_trigger_count = 0 # only MS triggered
  3230. co_trigger_count = 0 # both triggered
  3231. no_trigger_count = 0 # neither triggered
  3232. for i in range(len(trial_data[0])):
  3233. t_id = trial_data[0][i]
  3234. # Only process paired sessions (no 'M1' or 'M2' in t_id)
  3235. if 'M1' in t_id or 'M2' in t_id:
  3236. continue
  3237. m1_behavior = trial_data[4][i]
  3238. m2_behavior = trial_data[5][i]
  3239. m1_trigger_value = trial_data[6][i]
  3240. m2_trigger_value = trial_data[7][i]
  3241. group_num = t_id.split('B')[1].split('D')[0]
  3242. if group_num in md_is_m2_groups:
  3243. md_trigger_value = m2_trigger_value # M2 is MD
  3244. ms_trigger_value = m1_trigger_value # M1 is MS
  3245. md_behavior = m2_behavior
  3246. ms_behavior = m1_behavior
  3247. else:
  3248. md_trigger_value = m1_trigger_value # M1 is MD
  3249. ms_trigger_value = m2_trigger_value # M2 is MS
  3250. md_behavior = m1_behavior
  3251. ms_behavior = m2_behavior
  3252. md_triggered = (md_trigger_value != 0) and (md_behavior != 'N')
  3253. ms_triggered = (ms_trigger_value != 0) and (ms_behavior != 'N')
  3254. if md_triggered and ms_triggered:
  3255. co_trigger_count += 1
  3256. elif md_triggered:
  3257. md_trigger_count += 1
  3258. elif ms_triggered:
  3259. ms_trigger_count += 1
  3260. else:
  3261. no_trigger_count += 1
  3262. total_trials = md_trigger_count + ms_trigger_count + co_trigger_count + no_trigger_count
  3263. print(f"Paired session trial counts:")
  3264. print(f" md-only trigger: {md_trigger_count}")
  3265. print(f" ms-only trigger: {ms_trigger_count}")
  3266. print(f" co-trigger (both): {co_trigger_count}")
  3267. print(f" no trigger: {no_trigger_count}")
  3268. print(f" total: {total_trials}")
  3269. print(f" md-trigger total (md-only + co): {md_trigger_count + co_trigger_count}")
  3270. print(f" ms-trigger total (ms-only + co): {ms_trigger_count + co_trigger_count}")
  3271. df_co_trigger = pd.DataFrame({
  3272. 'category': ['md-only', 'ms-only', 'co-trigger'],
  3273. 'md_count': [md_trigger_count, 0, co_trigger_count],
  3274. 'ms_count': [0, ms_trigger_count, co_trigger_count]
  3275. })
  3276. df_co_trigger.to_csv('data/FigS1D_co_trigger.csv', index=False)
  3277. # %%
  3278. import numpy as np
  3279. import matplotlib.pyplot as plt
  3280. df = pd.read_csv('data/FigS1D_co_trigger.csv')
  3281. dom_vals = [df.loc[0, 'md_count'], df.loc[2, 'md_count']]
  3282. sub_vals = [df.loc[1, 'ms_count'], df.loc[2, 'ms_count']]
  3283. labels = ['single-trigger', 'co-trigger']
  3284. x = np.array([0, 1]) # Dom, Sub
  3285. width = 0.6
  3286. fig, ax = plt.subplots()
  3287. bottom_dom = 0
  3288. bottom_sub = 0
  3289. ax.bar(0, dom_vals[0], width, bottom=bottom_dom, color='blue')
  3290. ax.bar(0, dom_vals[1], width, bottom=dom_vals[0], color='orange')
  3291. ax.bar(1, sub_vals[0], width, bottom=bottom_sub, color='blue')
  3292. ax.bar(1, sub_vals[1], width, bottom=sub_vals[0], color='orange')
  3293. ax.set_xticks(x)
  3294. ax.set_xticklabels(['Dom', 'Sub'])
  3295. ax.set_ylim(0, 80)
  3296. ax.set_ylabel('Trials')
  3297. plt.tight_layout()
  3298. plt.show()
  3299. # %% [markdown]
  3300. # ### FigS1E. 3 type of grooming time (%)
  3301. # %%
  3302. from analyze_data_utils import filter_dict_data, filter_in_range, calculate_duration_time, calculate_total, calculate_duration_pct, dict_to_dataframe
  3303. self_grooming_data = filter_dict_data(lst_bhvr_labels_dict, 'grooming')
  3304. self_grooming_data = filter_in_range(self_grooming_data, [750,1950], method='replace')
  3305. self_grooming_durations = calculate_duration_time(self_grooming_data)
  3306. total_self_grooming_duration = calculate_total(self_grooming_durations)
  3307. self_grooming_pct = calculate_duration_pct(total_self_grooming_duration, time_range=[20, 60])
  3308. avg_self_grooming_duration_df = dict_to_dataframe(self_grooming_pct, value_name='self_grooming_pct', nan2zero=True)
  3309. # avg_self_grooming_duration_df.to_csv('data/fig1I_suppl_avg_self_grooming_duration.csv', index=False)
  3310. self_grooming_pct_means = avg_self_grooming_duration_df.mean(numeric_only=True)
  3311. grooming_partner_data = filter_dict_data(lst_bhvr_labels_dict, 'groom_partner')
  3312. grooming_partner_data = filter_in_range(grooming_partner_data, [750,1950], method='replace')
  3313. grooming_partner_durations = calculate_duration_time(grooming_partner_data)
  3314. total_grooming_partner_duration = calculate_total(grooming_partner_durations)
  3315. grooming_partner_pct = calculate_duration_pct(total_grooming_partner_duration, time_range=[20, 60])
  3316. avg_grooming_partner_duration_df = dict_to_dataframe(grooming_partner_pct, value_name='grooming_given_pct', nan2zero=True)
  3317. grooming_given_pct_means = avg_grooming_partner_duration_df.mean(numeric_only=True)
  3318. # Calculate grooming received percentages by swapping the given percentages
  3319. grooming_received_pct_means = pd.Series({
  3320. 'grooming_given_pct_sD': grooming_given_pct_means['grooming_given_pct_sS'],
  3321. 'grooming_given_pct_sS': grooming_given_pct_means['grooming_given_pct_sD'],
  3322. 'grooming_given_pct_pD': grooming_given_pct_means['grooming_given_pct_pS'],
  3323. 'grooming_given_pct_pS': grooming_given_pct_means['grooming_given_pct_pD']
  3324. })
  3325. grooming_summary = pd.DataFrame({
  3326. 'self-grooming': [
  3327. self_grooming_pct_means['self_grooming_pct_sD'],
  3328. self_grooming_pct_means['self_grooming_pct_pD'],
  3329. self_grooming_pct_means['self_grooming_pct_sS'],
  3330. self_grooming_pct_means['self_grooming_pct_pS']
  3331. ],
  3332. 'grooming-received': [
  3333. grooming_received_pct_means['grooming_given_pct_sD'],
  3334. grooming_received_pct_means['grooming_given_pct_pD'],
  3335. grooming_received_pct_means['grooming_given_pct_sS'],
  3336. grooming_received_pct_means['grooming_given_pct_pS']
  3337. ],
  3338. 'grooming-given': [
  3339. grooming_given_pct_means['grooming_given_pct_sD'],
  3340. grooming_given_pct_means['grooming_given_pct_pD'],
  3341. grooming_given_pct_means['grooming_given_pct_sS'],
  3342. grooming_given_pct_means['grooming_given_pct_pS']
  3343. ]
  3344. }, index=['sD', 'pD', 'sS', 'pS'])
  3345. grooming_summary.to_csv('data/FigS1E_grooming_3type.csv')
  3346. # %%
  3347. import pandas as pd
  3348. import numpy as np
  3349. import matplotlib.pyplot as plt
  3350. df = pd.read_csv('data/FigS1E_grooming_3type.csv')
  3351. df = df.set_index('Unnamed: 0').loc[['sD','pD','sS','pS']]
  3352. x_labels = ['sD','pD','sS','pS']
  3353. x = np.arange(len(x_labels))
  3354. width = 0.6
  3355. colors = {
  3356. 'self-grooming': 'blue',
  3357. 'grooming-received': 'green',
  3358. 'grooming-given': 'orange'
  3359. }
  3360. fig, ax = plt.subplots()
  3361. bottom = np.zeros(len(x_labels))
  3362. for col in ['self-grooming', 'grooming-received', 'grooming-given']:
  3363. ax.bar(x, df[col].values, width, bottom=bottom, color=colors[col])
  3364. bottom += df[col].values
  3365. ax.set_xticks(x)
  3366. ax.set_xticklabels(x_labels)
  3367. ax.set_ylim(0, 15)
  3368. ax.set_ylabel('Percentage')
  3369. plt.tight_layout()
  3370. plt.show()
  3371. # %% [markdown]
  3372. # ## Fig2
  3373. # %% [markdown]
  3374. # ### Fig2B&D. location and velocity histogram
  3375. # %%
  3376. import numpy as np
  3377. import matplotlib.pyplot as plt
  3378. from matplotlib.colors import ListedColormap, BoundaryNorm
  3379. import seaborn as sns
  3380. from analyze_data_utils import ret_extract_target_data
  3381. import pickle
  3382. import pandas as pd
  3383. def prepare_velocity_data(categories, groups, ranks, data_source, min_speed=0, max_speed=300):
  3384. all_x = []
  3385. all_y = []
  3386. all_vx = []
  3387. all_vy = []
  3388. all_ids = []
  3389. time = 'withrat'
  3390. x_offset = 12
  3391. for categorie in categories:
  3392. for group in groups:
  3393. for rank in ranks:
  3394. with open ("G:\Li_lab\ppt\S_paper\Paper_v3\check_ret_zone\zone_correct.pkl", 'rb') as f:
  3395. zone_correct = pickle.load(f)
  3396. m_id = group+rank
  3397. dx = (zone_correct['w_correct'][m_id] - 0.5) / 24
  3398. rat_x = zone_correct['x_correct'][m_id] - x_offset
  3399. x_min, x_max = rat_x, rat_x+(24+x_offset)*dx
  3400. y_min, y_max = -10.5, 9.5
  3401. interested_frame = [0, 8900]
  3402. x_center_bl = ret_extract_target_data(ret_raw_data, categories=[categorie],
  3403. groups=[group], ranks=[rank],
  3404. target_key=['X center'], times=[time])
  3405. y_center_bl = ret_extract_target_data(ret_raw_data, categories=[categorie],
  3406. groups=[group], ranks=[rank],
  3407. target_key=['Y center'], times=[time])
  3408. # velocity_al = ret_extract_target_data(ret_raw_data, categories=[categorie],
  3409. # groups=[group], ranks=[rank],
  3410. # target_key=['Velocity'], times=['withrat'])
  3411. # data_to_save = {'x_center_al': x_center_bl, 'y_center_al': y_center_bl, 'velocity_al': velocity_al}
  3412. # with open(r'G:\Li_lab\ppt\S_paper\Paper_v3\check_ret_zone\test_center_data.pkl', 'wb') as f:
  3413. # pickle.dump(data_to_save, f)
  3414. x = x_center_bl[categorie][rank][time][group].iloc[interested_frame[0]:interested_frame[1]]
  3415. y = y_center_bl[categorie][rank][time][group].iloc[interested_frame[0]:interested_frame[1]]
  3416. if m_id == 'D4MD':
  3417. x = x[~x.index.isin([7181, 7182])]
  3418. y = y[~y.index.isin([7181, 7182])]
  3419. x = pd.to_numeric(x, errors='coerce')
  3420. y = pd.to_numeric(y, errors='coerce')
  3421. x = np.array(x, dtype=float)
  3422. y = np.array(y, dtype=float)
  3423. vx = np.diff(x) * 30
  3424. vy = np.diff(y) * 30
  3425. speed = np.sqrt(vx**2 + vy**2)
  3426. mask = (speed >= min_speed) & (speed <= max_speed) & \
  3427. (x[:-1] >= x_min) & (x[:-1] < x_max) & \
  3428. (y[:-1] >= y_min) & (y[:-1] < y_max)
  3429. all_x.extend(x[:-1][mask])
  3430. all_y.extend(y[:-1][mask])
  3431. all_vx.extend(vx[mask])
  3432. all_vy.extend(vy[mask])
  3433. all_ids.extend([group+rank] * np.sum(mask))
  3434. # Convert to NumPy arrays
  3435. all_x = np.array(all_x)
  3436. all_y = np.array(all_y)
  3437. all_vx = np.array(all_vx)
  3438. all_vy = np.array(all_vy)
  3439. all_ids = np.array(all_ids)
  3440. return all_x, all_y, all_vx, all_vy, all_ids
  3441. groups = ['D1', 'D2', 'D3', 'D4', 'D5', 'D6', 'D7', 'D8', 'D9', 'D10', 'D11']
  3442. ranks = ['MD', 'MS']
  3443. xbins = 12
  3444. ybins = 10
  3445. x, y, vx, vy, point_ids = prepare_velocity_data(['single'], groups, ranks, ret_raw_data)
  3446. count, xedges, yedges = np.histogram2d(x, y, bins=[xbins, ybins])
  3447. mouse_count_grid = np.zeros_like(count, dtype=int)
  3448. x_bin_indices = np.digitize(x, xedges) - 1
  3449. y_bin_indices = np.digitize(y, yedges) - 1
  3450. cell_mice_dict = {}
  3451. for xi, yi, mouse_id in zip(x_bin_indices, y_bin_indices, point_ids):
  3452. if 0 <= xi < xbins and 0 <= yi < ybins:
  3453. key = (xi, yi)
  3454. if key not in cell_mice_dict:
  3455. cell_mice_dict[key] = set()
  3456. cell_mice_dict[key].add(mouse_id)
  3457. for (xi, yi), mice_set in cell_mice_dict.items():
  3458. mouse_count_grid[xi, yi] = len(mice_set)
  3459. mask = (mouse_count_grid >= 2) & (count >= 5)
  3460. grid_vx, _, _ = np.histogram2d(x, y, bins=[xedges, yedges], weights=vx)
  3461. grid_vy, _, _ = np.histogram2d(x, y, bins=[xedges, yedges], weights=vy)
  3462. grid_vx /= count
  3463. grid_vy /= count
  3464. grid_vx[~mask] = 1e-5
  3465. grid_vy[~mask] = 1e-5
  3466. speed_grid = np.sqrt(grid_vx**2 + grid_vy**2)
  3467. speeds = np.sqrt(vx**2 + vy**2)
  3468. speed_sum, _, _ = np.histogram2d(x, y, bins=[xedges, yedges], weights=speeds)
  3469. speed_avg = speed_sum / count
  3470. speed_avg[~mask] = 1e-5
  3471. heatmap, xedges, yedges = np.histogram2d(x, y, bins=[xedges, yedges])
  3472. hist_frequency = heatmap / np.sum(heatmap) * 100
  3473. hist_frequency[~mask] = 1e-5
  3474. df_grid = pd.DataFrame({
  3475. 'x_index': np.tile(np.arange(xbins), ybins),
  3476. 'y_index': np.repeat(np.arange(ybins), xbins),
  3477. 'time_pct': hist_frequency.T.flatten(),
  3478. 'speed_avg': speed_avg.T.flatten(),
  3479. 'vx': grid_vx.T.flatten(),
  3480. 'vy': grid_vy.T.flatten()
  3481. })
  3482. df_grid.to_csv('data/Fig2B&D_location_speed_heatmap.csv', index=False)
  3483. # %%
  3484. import numpy as np
  3485. import matplotlib.pyplot as plt
  3486. import pandas as pd
  3487. # Load data
  3488. df = pd.read_csv('data/Fig2B&D_location_speed_heatmap.csv')
  3489. xbins = 12
  3490. ybins = 10
  3491. # Reshape to matrices (shape: (ybins, xbins))
  3492. hist_frequency = df.pivot(index='y_index', columns='x_index', values='time_pct').values
  3493. speed_avg = df.pivot(index='y_index', columns='x_index', values='speed_avg').values
  3494. grid_vx = df.pivot(index='y_index', columns='x_index', values='vx').values
  3495. grid_vy = df.pivot(index='y_index', columns='x_index', values='vy').values
  3496. # Figure 1: Coordinate heatmap
  3497. fig, ax = plt.subplots(figsize=(5,5), dpi=200)
  3498. im = ax.imshow(
  3499. hist_frequency, # no transpose (shape: 10x12)
  3500. cmap='viridis',
  3501. origin='lower',
  3502. extent=(0, xbins, 0, ybins),
  3503. vmin=0, vmax=12
  3504. )
  3505. fig.colorbar(im, label='Time (%)')
  3506. ax.set_xlim(0, xbins)
  3507. ax.set_ylim(0, ybins)
  3508. ax.set_xticks([])
  3509. ax.set_yticks([])
  3510. ax.invert_yaxis()
  3511. plt.tight_layout()
  3512. # plt.savefig(...)
  3513. plt.show()
  3514. # Figure 2: Speed heatmap with arrows
  3515. fig, ax = plt.subplots(figsize=(5,5), dpi=200)
  3516. # Speed heatmap
  3517. im2 = ax.imshow(
  3518. speed_avg, # no transpose
  3519. cmap='viridis',
  3520. origin='lower',
  3521. extent=(0, xbins, 0, ybins),
  3522. vmin=0, vmax=20
  3523. )
  3524. plt.colorbar(im2, label='Speed (cm/s)')
  3525. # Arrows
  3526. X, Y = np.meshgrid(np.arange(xbins) + 0.5, np.arange(ybins) + 0.5)
  3527. scale = 0.3
  3528. head_scale = 0.1
  3529. for i in range(X.shape[0]): # y index (0 ~ ybins-1)
  3530. for j in range(X.shape[1]): # x index (0 ~ xbins-1)
  3531. vx_ = grid_vx[i, j]
  3532. vy_ = grid_vy[i, j]
  3533. mag = np.sqrt(vx_**2 + vy_**2)
  3534. dx = vx_ * scale
  3535. dy = vy_ * scale
  3536. x0 = X[i, j]
  3537. y0 = Y[i, j]
  3538. x1 = x0 + dx
  3539. y1 = y0 + dy
  3540. ax.annotate('', xy=(x1, y1), xytext=(x0, y0),
  3541. arrowprops=dict(
  3542. arrowstyle='->, head_width={:.2f}, head_length={:.2f}'.format(
  3543. mag * head_scale, mag * head_scale * 1.5
  3544. ),
  3545. color='white',
  3546. linewidth=1,
  3547. mutation_scale=5,
  3548. shrinkA=0, shrinkB=0
  3549. ))
  3550. ax.invert_yaxis()
  3551. ax.set_xlim(0, xbins)
  3552. ax.set_ylim(0, ybins)
  3553. ax.set_xticks([])
  3554. ax.set_yticks([])
  3555. plt.tight_layout()
  3556. # plt.savefig(...)
  3557. plt.show()
  3558. # %% [markdown]
  3559. # ### Fig2C&E & 2C suppl: Time in zone, Speed in zone (ret, withrat 0-8900)
  3560. # %%
  3561. import numpy as np
  3562. import pandas as pd
  3563. from analyze_data_utils import ret_extract_target_data, get_ret_location_value
  3564. # Define parameters
  3565. fps = 30
  3566. max_frames = 8900 # withrat 0-8900
  3567. # Arena parameters (consistent with other ret analyses)
  3568. dx = (19 + 5) / 24 # cm per unit (= 1.0)
  3569. rat_x = -5 # rat side X coordinate start (cm)
  3570. arena_size = 24 * dx # arena total length 24 cm
  3571. # Zone definitions: near:middle:far = 9:9:6 (cm)
  3572. near_zone_width = 9 * dx
  3573. middle_zone_width = 9 * dx
  3574. far_zone_width = 6 * dx
  3575. near_zone_range = (rat_x, rat_x + near_zone_width) # (-5, 4)
  3576. middle_zone_range = (rat_x + near_zone_width, rat_x + near_zone_width + middle_zone_width) # ( 4, 13)
  3577. far_zone_range = (rat_x + near_zone_width + middle_zone_width, rat_x + arena_size) # (13, 19)
  3578. zone_names = {0: 'near', 1: 'middle', 2: 'far'}
  3579. ret_groups = [f'D{i}' for i in range(1, 12)] # D1 ~ D11
  3580. # Extract X center and Y center (withrat period)
  3581. x_center_ret = ret_extract_target_data(ret_raw_data, target_key=['X center'], times=['withrat'])
  3582. y_center_ret = ret_extract_target_data(ret_raw_data, target_key=['Y center'], times=['withrat'])
  3583. # ret_zones_time_pcts[condition][rank][zone] = [value for D1..D11]
  3584. ret_zones_time_pcts = {}
  3585. ret_zones_speed_vals = {}
  3586. for cat in ['single', 'pair']:
  3587. ret_zones_time_pcts[cat] = {'MD': {0: [], 1: [], 2: []}, 'MS': {0: [], 1: [], 2: []}}
  3588. ret_zones_speed_vals[cat] = {'MD': {0: [], 1: [], 2: []}, 'MS': {0: [], 1: [], 2: []}}
  3589. for cat in ['single', 'pair']:
  3590. for rank in ['MD', 'MS']:
  3591. for d in ret_groups:
  3592. try:
  3593. x_series = x_center_ret[cat][rank]['withrat'][d]
  3594. y_series = y_center_ret[cat][rank]['withrat'][d]
  3595. except (KeyError, TypeError):
  3596. for zone in [0, 1, 2]:
  3597. ret_zones_time_pcts[cat][rank][zone].append(np.nan)
  3598. ret_zones_speed_vals[cat][rank][zone].append(np.nan)
  3599. continue
  3600. # Coordinate time correction: align behavior frames with coordinate frames using rat_in_frames and offset_frames
  3601. # Determine correct session_id based on condition and rank
  3602. if cat == 'pair':
  3603. sess_id = d + 'M1&M2'
  3604. else: # cat == 'single'
  3605. day_num = d.split('D')[1].split('M')[0] # Extract number, e.g., 'D4' -> '4'
  3606. if day_num == '4': # D4 special case
  3607. m_id = '1' if rank == 'MS' else '2'
  3608. else: # all other sessions
  3609. m_id = '1' if rank == 'MD' else '2'
  3610. sess_id = d + 'M' + m_id
  3611. # Calculate offset
  3612. corr_rat_in = rat_in_frames[sess_id] - offset_frames[sess_id]
  3613. x_raw = pd.to_numeric(x_series, errors='coerce')
  3614. y_raw = pd.to_numeric(y_series, errors='coerce')
  3615. # Handle negative offset: behavior data starts earlier than coordinate data
  3616. if corr_rat_in >= 0:
  3617. # Normal case: take max_frames frames starting from corr_rat_in
  3618. x = x_raw.iloc[corr_rat_in:corr_rat_in + max_frames].reset_index(drop=True)
  3619. y = y_raw.iloc[corr_rat_in:corr_rat_in + max_frames].reset_index(drop=True)
  3620. else:
  3621. # Negative offset: pad NaN at the beginning, then take data from 0
  3622. n_pad = -corr_rat_in # Number of frames to pad
  3623. n_data = max_frames - n_pad # Number of frames to take from coordinate data
  3624. # Create padded NaN series
  3625. pad_series = pd.Series([np.nan] * n_pad)
  3626. # Extract coordinate data (starting from 0, take n_data frames)
  3627. x_data = x_raw.iloc[:n_data]
  3628. y_data = y_raw.iloc[:n_data]
  3629. # Concatenate
  3630. x = pd.concat([pad_series, x_data], ignore_index=True)
  3631. y = pd.concat([pad_series, y_data], ignore_index=True)
  3632. # Compute zone labels
  3633. location_value = get_ret_location_value(x)
  3634. # Compute frame-by-frame speed (cm/s) from X/Y displacement
  3635. dx_pos = x.diff().fillna(0)
  3636. dy_pos = y.diff().fillna(0)
  3637. speed = np.sqrt(dx_pos**2 + dy_pos**2) * fps
  3638. for zone in [0, 1, 2]:
  3639. zone_mask = (location_value == zone)
  3640. zone_count = zone_mask.sum()
  3641. zone_pct = zone_count / max_frames * 100
  3642. ret_zones_time_pcts[cat][rank][zone].append(zone_pct)
  3643. zone_speed_mean = speed[zone_mask].mean() if zone_mask.any() else np.nan
  3644. ret_zones_speed_vals[cat][rank][zone].append(zone_speed_mean)
  3645. # %%
  3646. import pandas as pd
  3647. import numpy as np
  3648. zone_order = ['near', 'middle', 'far']
  3649. zone_id_map = {'near': 0, 'middle': 1, 'far': 2}
  3650. ret_groups = [f'D{i}' for i in range(1, 12)]
  3651. # -- Time percentage DataFrame -----------------------------------------
  3652. df_ret_zones_pct = {'single': None, 'pair': None}
  3653. df_ret_zones_unit = {'single': None, 'pair': None}
  3654. for cat in ['single', 'pair']:
  3655. data_pct = []
  3656. for rank in ['MD', 'MS']:
  3657. for i, d in enumerate(ret_groups):
  3658. row = {'D_id': d, 'rank': rank}
  3659. for zname in zone_order:
  3660. vals = ret_zones_time_pcts[cat][rank][zone_id_map[zname]]
  3661. row[f'{zname}_pct'] = vals[i] if i < len(vals) else np.nan
  3662. data_pct.append(row)
  3663. df_ret_zones_pct[cat] = pd.DataFrame(data_pct)
  3664. # Normalize: % -> s/(min·cm²)
  3665. # withrat duration (min)
  3666. withrat_time_min = max_frames / fps / 60
  3667. # Zone areas (cm²): arena width 24 cm
  3668. zone_areas = {'near': 9 * 24, 'middle': 9 * 24, 'far': 6 * 24}
  3669. df_ret_zones_unit[cat] = df_ret_zones_pct[cat].copy()
  3670. for zname in zone_order:
  3671. col = f'{zname}_pct'
  3672. # pct / 100 * total_s / time_min / area_cm2 = s/(min·cm²)
  3673. df_ret_zones_unit[cat][col] = (
  3674. df_ret_zones_pct[cat][col] / 100 * 60 / zone_areas[zname]
  3675. )
  3676. # Adjust column order
  3677. # Info columns
  3678. info_cols = ['D_id', 'rank']
  3679. # Data columns
  3680. data_cols = [f'{zone}_pct' for zone in zone_order]
  3681. for cat in ['single', 'pair']:
  3682. df_ret_zones_pct[cat] = df_ret_zones_pct[cat][info_cols + data_cols]
  3683. df_ret_zones_unit[cat] = df_ret_zones_unit[cat][info_cols + data_cols]
  3684. # Save DataFrame (only save single condition)
  3685. df_ret_zones_pct['single'].to_csv('data/FigS2C_zones_time_pcts.csv', index=False)
  3686. df_ret_zones_unit['single'].to_csv('data/Fig2C_zones_time_unit.csv', index=False)
  3687. # -- Speed DataFrame -----------------------------------------------
  3688. df_ret_zones_speed = {'single': None, 'pair': None}
  3689. for cat in ['single', 'pair']:
  3690. data_spd = []
  3691. for rank in ['MD', 'MS']:
  3692. for i, d in enumerate(ret_groups):
  3693. row = {'D_id': d, 'rank': rank}
  3694. for zname in zone_order:
  3695. vals = ret_zones_speed_vals[cat][rank][zone_id_map[zname]]
  3696. row[f'{zname}_speed'] = vals[i] if i < len(vals) else np.nan
  3697. data_spd.append(row)
  3698. df_ret_zones_speed[cat] = pd.DataFrame(data_spd)
  3699. # Adjust column order
  3700. # Info columns
  3701. info_cols = ['D_id', 'rank']
  3702. # Data columns
  3703. data_cols = [f'{zone}_speed' for zone in zone_order]
  3704. for cat in ['single', 'pair']:
  3705. df_ret_zones_speed[cat] = df_ret_zones_speed[cat][info_cols + data_cols]
  3706. # Save DataFrame (only save single condition)
  3707. df_ret_zones_speed['single'].to_csv('data/Fig2E_zones_speed.csv', index=False)
  3708. # %%
  3709. df = pd.read_csv('data/FigS2C_zones_time_pcts.csv')
  3710. print(df.columns.tolist())
  3711. exclude_cols = ['B_id', 'rank', 'mouse_id']
  3712. plot_cols = [col for col in df.columns if col not in exclude_cols]
  3713. means = df[plot_cols].mean()
  3714. plt.figure()
  3715. bar_colors = ['orange', 'blue', 'green']
  3716. bars = plt.bar(means.index, means.values, color=bar_colors[:len(means)])
  3717. df_melted = df[plot_cols]
  3718. sns.stripplot(
  3719. data=df_melted,
  3720. jitter=True,
  3721. color='gray',
  3722. size=5
  3723. )
  3724. plt.ylabel('Time in zone (%)')
  3725. plt.xticks(rotation=45, ha='right')
  3726. plt.gca().invert_xaxis()
  3727. plt.ylim(0, 100)
  3728. plt.show()
  3729. # %%
  3730. df = pd.read_csv('data/Fig2C_zones_time_unit.csv')
  3731. print(df.columns.tolist())
  3732. exclude_cols = ['B_id', 'rank', 'mouse_id']
  3733. plot_cols = [col for col in df.columns if col not in exclude_cols]
  3734. means = df[plot_cols].mean()
  3735. plt.figure()
  3736. bar_colors = ['orange', 'blue', 'green']
  3737. bars = plt.bar(means.index, means.values, color=bar_colors[:len(means)])
  3738. df_melted = df[plot_cols]
  3739. sns.stripplot(
  3740. data=df_melted,
  3741. jitter=True,
  3742. color='gray',
  3743. size=5
  3744. )
  3745. plt.ylabel('Time in zone (s/(min﹡cm²))')
  3746. plt.xticks(rotation=45, ha='right')
  3747. plt.gca().invert_xaxis()
  3748. plt.ylim(0, 0.4)
  3749. plt.show()
  3750. # %%
  3751. df = pd.read_csv('data/Fig2E_zones_speed.csv')
  3752. print(df.columns.tolist())
  3753. exclude_cols = ['B_id', 'rank', 'mouse_id']
  3754. plot_cols = [col for col in df.columns if col not in exclude_cols]
  3755. means = df[plot_cols].mean()
  3756. plt.figure()
  3757. bar_colors = ['orange', 'blue', 'green']
  3758. bars = plt.bar(means.index, means.values, color=bar_colors[:len(means)])
  3759. df_melted = df[plot_cols]
  3760. sns.stripplot(
  3761. data=df_melted,
  3762. jitter=True,
  3763. color='gray',
  3764. size=5
  3765. )
  3766. plt.ylabel('Speed in zone (cm/s)')
  3767. plt.xticks(rotation=45, ha='right')
  3768. plt.gca().invert_xaxis()
  3769. plt.ylim(0, 9)
  3770. plt.show()
  3771. # %% [markdown]
  3772. # ### Fig2F. ΔTime in zone (%)
  3773. # %%
  3774. import pandas as pd
  3775. import numpy as np
  3776. # Define zone order
  3777. zone_order = ['near', 'middle', 'far']
  3778. ret_groups = [f'D{i}' for i in range(1, 12)]
  3779. delta_data = []
  3780. for d_id in ret_groups:
  3781. row = {'D_id': d_id}
  3782. for rank in ['MD', 'MS']:
  3783. single_row = df_ret_zones_pct['single'][
  3784. (df_ret_zones_pct['single']['D_id'] == d_id) &
  3785. (df_ret_zones_pct['single']['rank'] == rank)
  3786. ]
  3787. pair_row = df_ret_zones_pct['pair'][
  3788. (df_ret_zones_pct['pair']['D_id'] == d_id) &
  3789. (df_ret_zones_pct['pair']['rank'] == rank)
  3790. ]
  3791. for zone in zone_order:
  3792. col = f'{zone}_pct'
  3793. if len(single_row) > 0 and len(pair_row) > 0:
  3794. row[f'{zone}_{rank}'] = pair_row[col].values[0] - single_row[col].values[0]
  3795. else:
  3796. row[f'{zone}_{rank}'] = np.nan
  3797. delta_data.append(row)
  3798. # Create DataFrame and adjust column order
  3799. df_delta_time = pd.DataFrame(delta_data)
  3800. df_delta_time = df_delta_time[['D_id', 'near_MD', 'near_MS', 'middle_MD', 'middle_MS', 'far_MD', 'far_MS']]
  3801. df_delta_time.columns = [col.replace('_MD', '_D').replace('_MS', '_S') for col in df_delta_time.columns]
  3802. # Save result
  3803. df_delta_time.to_csv('data/Fig2F_delta_zones_time.csv', index=False)
  3804. # %%
  3805. df = pd.read_csv(f'data/Fig2F_delta_zones_time.csv')
  3806. print(df.columns.tolist())
  3807. groups = ['far', 'middle', 'near']
  3808. fig, ax = plt.subplots()
  3809. x_pos = {}
  3810. xticks = []
  3811. xticklabels = []
  3812. x = 0
  3813. colors = {'D': 'orange', 'S': 'blue'}
  3814. # 1. x
  3815. for g in groups:
  3816. for cond in ['D', 'S']:
  3817. if cond == 'D':
  3818. x_pos[f'{g}_{cond}_mean'] = x
  3819. x_pos[f'{g}_{cond}_pts'] = x + 0.8
  3820. xticks += [x, x + 0.8]
  3821. xticklabels += [f'{g}_{cond}_mean', f'{g}_{cond}_pts']
  3822. else:
  3823. x_pos[f'{g}_{cond}_pts'] = x
  3824. x_pos[f'{g}_{cond}_mean'] = x + 0.8
  3825. xticks += [x, x + 0.8]
  3826. xticklabels += [f'{g}_{cond}_pts', f'{g}_{cond}_mean']
  3827. x += 2.2
  3828. # 2. mean ± SD
  3829. for g in groups:
  3830. for cond in ['D', 'S']:
  3831. col = f'{g}_{cond}'
  3832. xpos = x_pos[f'{g}_{cond}_mean']
  3833. ax.errorbar(
  3834. xpos,
  3835. df[col].mean(),
  3836. yerr=df[col].std(),
  3837. fmt='o',
  3838. color=colors[cond],
  3839. capsize=4,
  3840. markersize=9, # 固定大小
  3841. elinewidth=1.5,
  3842. capthick=1.5,
  3843. zorder=3
  3844. )
  3845. # 3. raw points + pairing
  3846. for i in range(len(df)):
  3847. for g in groups:
  3848. D_col = f'{g}_D'
  3849. S_col = f'{g}_S'
  3850. x_d = x_pos[f'{g}_D_pts']
  3851. y_d = df.loc[i, D_col]
  3852. x_s = x_pos[f'{g}_S_pts']
  3853. y_s = df.loc[i, S_col]
  3854. ax.scatter(x_d, y_d, color=colors['D'])
  3855. ax.scatter(x_s, y_s, color=colors['S'])
  3856. ax.plot([x_d, x_s], [y_d, y_s],
  3857. color='gray', linewidth=1)
  3858. # 4. axis
  3859. ax.set_xticks(xticks)
  3860. ax.set_xticklabels(xticklabels, rotation=45, ha='right')
  3861. ax.set_ylabel('Δ Time in zone (%)')
  3862. ax.set_ylim([-60, 40])
  3863. plt.tight_layout()
  3864. plt.show()
  3865. # %% [markdown]
  3866. # ### Fig2G. Δ Speed in zone (cm/s)
  3867. # %%
  3868. import pandas as pd
  3869. import numpy as np
  3870. # Define zone order
  3871. zone_order = ['near', 'middle', 'far']
  3872. ret_groups = [f'D{i}' for i in range(1, 12)]
  3873. delta_data = []
  3874. for d_id in ret_groups:
  3875. row = {'D_id': d_id}
  3876. for rank in ['MD', 'MS']:
  3877. single_row = df_ret_zones_speed['single'][
  3878. (df_ret_zones_speed['single']['D_id'] == d_id) &
  3879. (df_ret_zones_speed['single']['rank'] == rank)
  3880. ]
  3881. pair_row = df_ret_zones_speed['pair'][
  3882. (df_ret_zones_speed['pair']['D_id'] == d_id) &
  3883. (df_ret_zones_speed['pair']['rank'] == rank)
  3884. ]
  3885. for zone in zone_order:
  3886. col = f'{zone}_speed'
  3887. if len(single_row) > 0 and len(pair_row) > 0:
  3888. row[f'{zone}_{rank}'] = pair_row[col].values[0] - single_row[col].values[0]
  3889. else:
  3890. row[f'{zone}_{rank}'] = np.nan
  3891. delta_data.append(row)
  3892. # Create DataFrame and adjust column order
  3893. df_delta_speed = pd.DataFrame(delta_data)
  3894. df_delta_speed = df_delta_speed[['D_id', 'near_MD', 'near_MS', 'middle_MD', 'middle_MS', 'far_MD', 'far_MS']]
  3895. df_delta_speed.columns = [col.replace('_MD', '_D').replace('_MS', '_S') for col in df_delta_speed.columns]
  3896. df_delta_speed.to_csv('data/Fig2G_delta_zones_speed.csv', index=False)
  3897. # %%
  3898. df = pd.read_csv(f'data/Fig2G_delta_zones_speed.csv')
  3899. print(df.columns.tolist())
  3900. groups = ['far', 'middle', 'near']
  3901. fig, ax = plt.subplots()
  3902. x_pos = {}
  3903. xticks = []
  3904. xticklabels = []
  3905. x = 0
  3906. colors = {'D': 'orange', 'S': 'blue'}
  3907. # 1. x
  3908. for g in groups:
  3909. for cond in ['D', 'S']:
  3910. if cond == 'D':
  3911. x_pos[f'{g}_{cond}_mean'] = x
  3912. x_pos[f'{g}_{cond}_pts'] = x + 0.8
  3913. xticks += [x, x + 0.8]
  3914. xticklabels += [f'{g}_{cond}_mean', f'{g}_{cond}_pts']
  3915. else:
  3916. x_pos[f'{g}_{cond}_pts'] = x
  3917. x_pos[f'{g}_{cond}_mean'] = x + 0.8
  3918. xticks += [x, x + 0.8]
  3919. xticklabels += [f'{g}_{cond}_pts', f'{g}_{cond}_mean']
  3920. x += 2.2
  3921. # 2. mean ± SD
  3922. for g in groups:
  3923. for cond in ['D', 'S']:
  3924. col = f'{g}_{cond}'
  3925. xpos = x_pos[f'{g}_{cond}_mean']
  3926. ax.errorbar(
  3927. xpos,
  3928. df[col].mean(),
  3929. yerr=df[col].std(),
  3930. fmt='o',
  3931. color=colors[cond],
  3932. capsize=4,
  3933. markersize=9,
  3934. elinewidth=1.5,
  3935. capthick=1.5,
  3936. zorder=3
  3937. )
  3938. # 3. raw points + pairing
  3939. for i in range(len(df)):
  3940. for g in groups:
  3941. D_col = f'{g}_D'
  3942. S_col = f'{g}_S'
  3943. x_d = x_pos[f'{g}_D_pts']
  3944. y_d = df.loc[i, D_col]
  3945. x_s = x_pos[f'{g}_S_pts']
  3946. y_s = df.loc[i, S_col]
  3947. ax.scatter(x_d, y_d, color=colors['D'])
  3948. ax.scatter(x_s, y_s, color=colors['S'])
  3949. ax.plot([x_d, x_s], [y_d, y_s],
  3950. color='gray', linewidth=1)
  3951. # 4. axis
  3952. ax.set_xticks(xticks)
  3953. ax.set_xticklabels(xticklabels, rotation=45, ha='right')
  3954. ax.set_ylabel('Δ Speed in zone (cm/s)')
  3955. ax.set_ylim([-4, 6])
  3956. plt.tight_layout()
  3957. plt.show()
  3958. # %% [markdown]
  3959. # ### Fig2H. with rat ethogram example
  3960. # %%
  3961. import numpy as np
  3962. import pandas as pd
  3963. import matplotlib.pyplot as plt
  3964. import matplotlib.colors as mcolors
  3965. import colorsys
  3966. from matplotlib.patches import FancyArrow
  3967. from analyze_data_utils import filter_in_range
  3968. def adjust_saturation(color, saturation):
  3969. rgb = mcolors.to_rgb(color)
  3970. h, l, s = colorsys.rgb_to_hls(*rgb)
  3971. new_s = saturation * s
  3972. new_rgb = colorsys.hls_to_rgb(h, l, new_s)
  3973. return new_rgb
  3974. with open('data/ret_etho_dict.pkl', 'rb') as f:
  3975. ret_etho_dict = pickle.load(f)
  3976. # ret_etho_dict = {
  3977. # 'single': {
  3978. # 'D6MD': ret_bhvr_labels_dict['s']['D']['D6MD'],
  3979. # 'D6MS': ret_bhvr_labels_dict['s']['S']['D6MS'],
  3980. # 'D7MD': ret_bhvr_labels_dict['s']['D']['D7MD'],
  3981. # 'D7MS': ret_bhvr_labels_dict['s']['S']['D7MS'],
  3982. # 'D8MD': ret_bhvr_labels_dict['s']['D']['D8MD'],
  3983. # 'D8MS': ret_bhvr_labels_dict['s']['S']['D8MS']
  3984. # },
  3985. # 'pair': {
  3986. # 'D6MD': ret_bhvr_labels_dict['p']['D']['D6MD'],
  3987. # 'D6MS': ret_bhvr_labels_dict['p']['S']['D6MS'],
  3988. # 'D7MD': ret_bhvr_labels_dict['p']['D']['D7MD'],
  3989. # 'D7MS': ret_bhvr_labels_dict['p']['S']['D7MS'],
  3990. # 'D8MD': ret_bhvr_labels_dict['p']['D']['D8MD'],
  3991. # 'D8MS': ret_bhvr_labels_dict['p']['S']['D8MS']
  3992. # }}
  3993. # with open('data/ret_etho_dict.pkl', 'wb') as f:
  3994. # pickle.dump(ret_etho_dict, f)
  3995. ret_framerange = [0, 8900]
  3996. saturation = 0.9
  3997. framerate=30
  3998. stim_frame = int(0.725 * framerate)
  3999. behavior_frames_dict = ret_etho_dict
  4000. behavior_frames_dict = dict(sorted(behavior_frames_dict.items()))
  4001. ret_behavior_frames_dict = filter_in_range(ret_etho_dict, ret_framerange, method='replace')
  4002. # %%
  4003. ret_behavior_properties = {
  4004. 'approach_partner': ('#143FCA', 4),
  4005. 'follow_partner': ('#137CAB', 4),
  4006. 'groom_partner': ('#0990FF', 4),
  4007. 'sniff_partner': ('#09FFFF', 4),
  4008. 'huddling': ('#6D57F3', 4),
  4009. 'approach': ('#FFCA09', 3),
  4010. 'dwelling': ('#FF8409', 3),
  4011. 'withdrawal': ('#FF0990', 3),
  4012. 'stretch_attend': ('#FF7979', 3),
  4013. 'freezing': ('#FF75EF', 3),
  4014. 'tail_rattling': ('#FF09FF', 3),
  4015. "rearing": ("#C409FF", 2),
  4016. 'grooming': ('#28AE61', 2),
  4017. "sniffing": ("#0FD400", 2),
  4018. 'in_proximity': ('white', 1),
  4019. 'rat_in': ('white', 1),
  4020. 'others': ('white', 1)}
  4021. saturation = 0.9
  4022. framerate = 30
  4023. fig, axes = plt.subplots(1, 2, figsize=(10, 2), dpi=300)
  4024. plt.subplots_adjust(wspace=0)
  4025. for i, group in enumerate(ret_etho_dict.keys()):
  4026. ax = axes[i]
  4027. ax.set_yticks(range(len(ret_etho_dict[group].keys())))
  4028. for idx, t_id in enumerate(ret_etho_dict[group].keys()):
  4029. behaviors = ret_etho_dict[group][t_id]
  4030. for b, frame_ranges in behaviors.items():
  4031. for frame_range in frame_ranges:
  4032. if np.isnan(frame_range).all():
  4033. continue
  4034. start_frame, end_frame = frame_range
  4035. color, zorder = ret_behavior_properties.get(b, ('white', 0))
  4036. color = adjust_saturation(color, saturation)
  4037. rect = plt.Rectangle((start_frame, idx - 0.4), end_frame - start_frame, 0.8, facecolor=color, edgecolor='none', zorder=zorder)
  4038. ax.add_patch(rect)
  4039. if idx % 2 == 0:
  4040. ax.axhline(idx - 0.5, color='black', linestyle='-', linewidth=1)
  4041. ytick_color = 'red'
  4042. else:
  4043. ax.axhline(idx - 0.5, color='black', linestyle='--', linewidth=0.1)
  4044. ytick_color = 'green'
  4045. ax.get_yticklabels()[idx].set_color(ytick_color)
  4046. if i == 0:
  4047. ax.set_yticklabels(['D', 'S', 'D', 'S', 'D', 'S'])
  4048. ax.set_ylabel('Rat')
  4049. else:
  4050. ax.set_yticks([])
  4051. ax.set_xticks([])
  4052. ret_mod_width = (ret_framerange[1] - ret_framerange[0]) / 12
  4053. ret_blank_width = (ret_framerange[1] - ret_framerange[0] + ret_mod_width) / 20
  4054. per_rect = plt.Rectangle((ret_framerange[0] - ret_mod_width - ret_blank_width, -0.5), ret_blank_width, 6,
  4055. facecolor='white', edgecolor='white', zorder=8)
  4056. ax.add_patch(per_rect)
  4057. post_rect = plt.Rectangle((ret_framerange[1], -0.5), ret_blank_width, 6,
  4058. facecolor='white', edgecolor='white', zorder=8)
  4059. ax.add_patch(post_rect)
  4060. ax.set_xlim(ret_framerange[0] - ret_blank_width - ret_mod_width, ret_framerange[1] + ret_blank_width)
  4061. ax.set_ylim(-0.5, 5.5)
  4062. ax.invert_yaxis()
  4063. ret_arrow = FancyArrow(0, -1.1, 0, 0.5, width=ret_blank_width/50, head_width=ret_blank_width/10, head_length=0.3, length_includes_head=True, color='red')
  4064. ret_arrow.set_clip_on(False)
  4065. ax.add_patch(ret_arrow)
  4066. threat_line = plt.Line2D([0, 8900], [5.75, 5.75], color='red', linewidth=2)
  4067. threat_line.set_clip_on(False)
  4068. ax.add_artist(threat_line)
  4069. if i == 1:
  4070. line_x1 = ret_framerange[1] - 30*60
  4071. line_x2 = ret_framerange[1]
  4072. ax.text((line_x1 + line_x2) / 2, -0.8, '60 sec', fontsize=8, color='black', ha='center')
  4073. ret_line = plt.Line2D([line_x1, line_x2], [-0.65, -0.65], color='black', linewidth=1)
  4074. ret_line.set_clip_on(False)
  4075. ax.add_artist(ret_line)
  4076. ret_legend_ = {
  4077. 'approach': ('#FFCA09', 3),
  4078. 'investigation': ('#FF8409', 3),
  4079. 'withdrawal': ('#FF0990', 3),
  4080. 'stretch-attend': ('#FF7979', 3),
  4081. 'freezing': ('#FF75EF', 3),
  4082. 'tail rattling': ('#FF09FF', 3),
  4083. "rearing": ("#C409FF", 2),
  4084. 'huddling': ('#6D57F3', 1),
  4085. 'approaching P': ('#143FCA', 4),
  4086. 'following P': ('#137CAB', 4),
  4087. 'grooming P': ('#0990FF', 4),
  4088. 'sniffing P': ('#09FFFF', 4),
  4089. 'grooming': ('#28AE61', 2),
  4090. "sniffing": ("#0FD400", 2),
  4091. 'other behaviors': ('white', 1)}
  4092. ret_legend = []
  4093. for label, (color, zorder) in ret_legend_.items():
  4094. if label == 'other behaviors':
  4095. rect = plt.Rectangle((0, 0), 1, 1, facecolor=adjust_saturation(color, saturation),
  4096. edgecolor='black', linewidth=0.5, label=label)
  4097. else:
  4098. rect = plt.Rectangle((0, 0), 1, 1, facecolor=adjust_saturation(color, saturation), label=label)
  4099. ret_legend.append(rect)
  4100. axes[1].legend(handles=ret_legend, loc='upper center', bbox_to_anchor=(0, -0.2), fontsize=7, ncol=7)
  4101. # plt.savefig('fig/fig_2/ret_behavior_ethogram.eps', format="eps", dpi=300, bbox_inches="tight")
  4102. plt.show()
  4103. # %%
  4104. # Assume ret_framerange = (0, 8900) # Infer from your actual situation, from original code
  4105. ret_framerange = (0, 8900) # If not explicitly defined, add this line
  4106. with open('data/ret_bl_etho_dict.pkl', 'rb') as f:
  4107. ret_bl_etho_dict = pickle.load(f)
  4108. # 1. Calculate ret_mod_width and ret_blank_width
  4109. ret_mod_width = (ret_framerange[1] - ret_framerange[0]) / 12 # ≈ 741.67
  4110. ret_blank_width = (ret_framerange[1] - ret_framerange[0] + ret_mod_width) / 20
  4111. # 2. Define baseline window (last ret_mod_width frames)
  4112. baseline_end = 8900 # consistent with ret_framerange[1]
  4113. window_start = baseline_end - ret_mod_width
  4114. window_end = baseline_end
  4115. # 3. Mapping function: baseline frame f -> transition zone new coordinate new_x
  4116. def map_to_transition(f):
  4117. # Map window_end (8900) to ret_framerange[0] (0)
  4118. # Map window_start to ret_framerange[0] - ret_mod_width
  4119. return ret_framerange[0] - ret_mod_width + (f - window_start)
  4120. # Start plotting
  4121. fig, axes = plt.subplots(1, 2, figsize=(10, 2), dpi=300)
  4122. plt.subplots_adjust(wspace=0)
  4123. for i, group in enumerate(ret_etho_dict.keys()):
  4124. ax = axes[i]
  4125. n_animals = len(ret_etho_dict[group])
  4126. ax.set_yticks(range(n_animals))
  4127. # ========== Step 1: Draw baseline transition data (left) ==========
  4128. # Assume ret_bl_etho_dict has the same group keys
  4129. if group in ret_bl_etho_dict:
  4130. bl_group_data = ret_bl_etho_dict[group]
  4131. for idx, t_id in enumerate(bl_group_data.keys()):
  4132. behaviors = bl_group_data[t_id]
  4133. for b, frame_ranges in behaviors.items():
  4134. for frame_range in frame_ranges:
  4135. # Skip invalid values (nan or None)
  4136. try:
  4137. if np.isnan(frame_range).all():
  4138. continue
  4139. except:
  4140. continue
  4141. start_frame, end_frame = frame_range
  4142. # Check overlap with transition window [window_start, window_end]
  4143. if end_frame <= window_start or start_frame >= window_end:
  4144. continue
  4145. # Clip overlap
  4146. clip_start = max(start_frame, window_start)
  4147. clip_end = min(end_frame, window_end)
  4148. # Map to new coordinates
  4149. new_start = map_to_transition(clip_start)
  4150. new_end = map_to_transition(clip_end)
  4151. width = new_end - new_start
  4152. if width <= 0:
  4153. continue
  4154. color, zorder = ret_behavior_properties.get(b, ('white', 0))
  4155. color = adjust_saturation(color, saturation)
  4156. rect = plt.Rectangle((new_start, idx - 0.4), width, 0.8,
  4157. facecolor=color, edgecolor='none', zorder=zorder)
  4158. ax.add_patch(rect)
  4159. # ========== Step 2: Draw original ret_etho_dict data (right) ==========
  4160. for idx, t_id in enumerate(ret_etho_dict[group].keys()):
  4161. behaviors = ret_etho_dict[group][t_id]
  4162. for b, frame_ranges in behaviors.items():
  4163. for frame_range in frame_ranges:
  4164. if np.isnan(frame_range).all():
  4165. continue
  4166. start_frame, end_frame = frame_range
  4167. color, zorder = ret_behavior_properties.get(b, ('white', 0))
  4168. color = adjust_saturation(color, saturation)
  4169. rect = plt.Rectangle((start_frame, idx - 0.4), end_frame - start_frame, 0.8,
  4170. facecolor=color, edgecolor='none', zorder=zorder)
  4171. ax.add_patch(rect)
  4172. # Draw separator lines between animals
  4173. if idx % 2 == 0:
  4174. ax.axhline(idx - 0.5, color='black', linestyle='-', linewidth=1)
  4175. ytick_color = 'red'
  4176. else:
  4177. ax.axhline(idx - 0.5, color='black', linestyle='--', linewidth=0.1)
  4178. ytick_color = 'green'
  4179. # Set ytick color (handled outside loop, we'll set later)
  4180. # Here only draw lines, not labels
  4181. # Set y-axis labels
  4182. if i == 0:
  4183. # Dynamically generate labels based on animal count (assuming D/S alternating)
  4184. labels = ['D' if j % 2 == 0 else 'S' for j in range(n_animals)]
  4185. ax.set_yticklabels(labels)
  4186. ax.set_ylabel('Rat')
  4187. # Set label colors based on parity
  4188. for j, label in enumerate(ax.get_yticklabels()):
  4189. label.set_color('red' if j % 2 == 0 else 'green')
  4190. else:
  4191. ax.set_yticks([])
  4192. # ========== Step 3: Set x-axis range and right blank rectangle ==========
  4193. ax.set_xticks([])
  4194. # Left white rectangle (per_rect) — keep unchanged!
  4195. per_rect = plt.Rectangle((ret_framerange[0] - ret_mod_width - ret_blank_width, -0.5),
  4196. ret_blank_width, n_animals,
  4197. facecolor='white', edgecolor='white', zorder=8)
  4198. ax.add_patch(per_rect)
  4199. # Right white rectangle (post_rect)
  4200. post_rect = plt.Rectangle((ret_framerange[1], -0.5),
  4201. ret_blank_width, n_animals,
  4202. facecolor='white', edgecolor='white', zorder=8)
  4203. ax.add_patch(post_rect)
  4204. # Set x-axis range: include per_rect, transition, main data, post_rect
  4205. ax.set_xlim(ret_framerange[0] - ret_mod_width - ret_blank_width,
  4206. ret_framerange[1] + ret_blank_width)
  4207. ax.set_ylim(-0.5, n_animals - 0.5)
  4208. ax.invert_yaxis()
  4209. # Add arrow (left of transition zone)
  4210. ret_arrow = FancyArrow(0, -1.1, 0, 0.5, width=ret_blank_width/50, head_width=ret_blank_width/10, head_length=0.3, length_includes_head=True, color='red')
  4211. ret_arrow.set_clip_on(False)
  4212. ax.add_patch(ret_arrow)
  4213. # Threat line (using actual x-axis range)
  4214. threat_line = plt.Line2D([0, 8900], [5.75, 5.75], color='red', linewidth=2)
  4215. threat_line.set_clip_on(False)
  4216. ax.add_artist(threat_line)
  4217. # Time scale (only right subplot)
  4218. if i == 1:
  4219. line_x1 = ret_framerange[1] - 30*60
  4220. line_x2 = ret_framerange[1]
  4221. ax.text((line_x1 + line_x2) / 2, -0.8, '60 sec', fontsize=8, color='black', ha='center')
  4222. ret_line = plt.Line2D([line_x1, line_x2], [-0.65, -0.65], color='black', linewidth=1)
  4223. ret_line.set_clip_on(False)
  4224. ax.add_artist(ret_line)
  4225. # Legend remains unchanged
  4226. ret_legend = []
  4227. for label, (color, zorder) in ret_legend_.items():
  4228. if label == 'other behaviors':
  4229. rect = plt.Rectangle((0, 0), 1, 1, facecolor=adjust_saturation(color, saturation),
  4230. edgecolor='black', linewidth=0.5, label=label)
  4231. else:
  4232. rect = plt.Rectangle((0, 0), 1, 1, facecolor=adjust_saturation(color, saturation), label=label)
  4233. ret_legend.append(rect)
  4234. axes[1].legend(handles=ret_legend, loc='upper center', bbox_to_anchor=(0, -0.2), fontsize=7, ncol=7)
  4235. # plt.savefig('fig/fig_2/ret_behavior_ethogram.eps', format="eps", dpi=300, bbox_inches="tight")
  4236. plt.show()
  4237. # %% [markdown]
  4238. # ### Fig2I. Social time(s/(min∙cm2))
  4239. # %%
  4240. import copy
  4241. import numpy as np
  4242. import pandas as pd
  4243. from analyze_data_utils import ret_extract_target_data, filter_dict_data, filter_in_range, merge_dicts, dict_to_dataframe, calculate_duration_time, calculate_total
  4244. # Define function to filter social behaviors within far zone
  4245. def filter_social_in_far(social_frames, location_series):
  4246. """
  4247. social_frames: [(start, end), ...]
  4248. location_series: pd.Series, True if in far zone
  4249. Returns: [(far_start, far_end), ...]
  4250. """
  4251. social_in_far = []
  4252. for start, end in social_frames:
  4253. loc_segment = location_series[start:end+1]
  4254. in_far_mask = loc_segment
  4255. if not in_far_mask.any():
  4256. continue
  4257. in_far_indices = loc_segment.index[in_far_mask]
  4258. group_start = None
  4259. for idx in in_far_indices:
  4260. if group_start is None:
  4261. group_start = idx
  4262. prev_idx = idx
  4263. elif idx == prev_idx + 1:
  4264. prev_idx = idx
  4265. else:
  4266. social_in_far.append((group_start, prev_idx))
  4267. group_start = idx
  4268. prev_idx = idx
  4269. if group_start is not None:
  4270. social_in_far.append((group_start, prev_idx))
  4271. if social_in_far == []:
  4272. social_in_far = [np.nan]
  4273. return social_in_far
  4274. # Merge social behavior data
  4275. ap_bl_data = filter_dict_data(ret_origin_bhvr_labels_dict, 'approach_partner')
  4276. sp_bl_data = filter_dict_data(ret_origin_bhvr_labels_dict, 'sniff_partner')
  4277. fp_bl_data = filter_dict_data(ret_origin_bhvr_labels_dict, 'follow_partner')
  4278. gp_bl_data = filter_dict_data(ret_origin_bhvr_labels_dict, 'groom_partner')
  4279. hp_bl_data = filter_dict_data(ret_origin_bhvr_labels_dict, 'huddling')
  4280. merged_bl_data = merge_dicts(ap_bl_data, sp_bl_data)
  4281. merged_bl_data = merge_dicts(merged_bl_data, fp_bl_data)
  4282. merged_bl_data = merge_dicts(merged_bl_data, gp_bl_data)
  4283. merged_bl_data = merge_dicts(merged_bl_data, hp_bl_data)
  4284. ap_wr_data = filter_dict_data(ret_bhvr_labels_dict, 'approach_partner')
  4285. sp_wr_data = filter_dict_data(ret_bhvr_labels_dict, 'sniff_partner')
  4286. fp_wr_data = filter_dict_data(ret_bhvr_labels_dict, 'follow_partner')
  4287. gp_wr_data = filter_dict_data(ret_bhvr_labels_dict, 'groom_partner')
  4288. hp_wr_data = filter_dict_data(ret_bhvr_labels_dict, 'huddling')
  4289. merged_wr_data = merge_dicts(ap_wr_data, sp_wr_data)
  4290. merged_wr_data = merge_dicts(merged_wr_data, fp_wr_data)
  4291. merged_wr_data = merge_dicts(merged_wr_data, gp_wr_data)
  4292. merged_wr_data = merge_dicts(merged_wr_data, hp_wr_data)
  4293. merged_bl_data = filter_in_range(merged_bl_data, [0,8900], method='replace')
  4294. merged_wr_data = filter_in_range(merged_wr_data, [0,8900], method='replace')
  4295. merged_data = {'s': merged_bl_data['p'], 'p': merged_wr_data['p']}
  4296. # Extract position data and filter social behaviors within far zone
  4297. dx = (19 + 5) / 24
  4298. rat_x = -5
  4299. max_frames = 8900
  4300. x_center_al = ret_extract_target_data(ret_raw_data, target_key=['X center'], times=['withrat'])
  4301. social_in_far_data = {}
  4302. for m in ['D', 'S']:
  4303. cond = 'MD' if m == 'D' else 'MS'
  4304. for d in range(1, 12):
  4305. session = f'D{d}{cond}'
  4306. series_x = x_center_al['pair'][cond]['withrat'][f'D{d}']
  4307. s = pd.to_numeric(series_x.iloc[:max_frames], errors='coerce')
  4308. far_mask = ((s >= rat_x + 18 * dx) & (s <= rat_x + 24 * dx)).fillna(False)
  4309. far_series = pd.Series(far_mask.values, index=range(len(far_mask)))
  4310. social_frames = merged_data.get('p', {}).get(m, {}).get(session, [])
  4311. if m not in social_in_far_data:
  4312. social_in_far_data[m] = {}
  4313. if not np.isnan(social_frames).any():
  4314. social_in_far = filter_social_in_far(social_frames, far_series)
  4315. social_in_far_data[m][session] = social_in_far
  4316. else:
  4317. social_in_far_data[m][session] = [np.nan]
  4318. # Calculate duration
  4319. social_in_far_data_dict = {'p': social_in_far_data}
  4320. social_in_far_durations = calculate_duration_time(social_in_far_data_dict)
  4321. total_social_in_far_duration = calculate_total(social_in_far_durations)
  4322. total_social_in_far_duration_df = dict_to_dataframe(total_social_in_far_duration, value_name='in_far_duration', groups=['p'], nan2zero=True)
  4323. raw_social_data = copy.deepcopy(merged_data)
  4324. social_data = {'p': raw_social_data['p']}
  4325. social_durations = calculate_duration_time(social_data)
  4326. total_social_duration = calculate_total(social_durations)
  4327. total_social_duration_df = dict_to_dataframe(total_social_duration, value_name='total_social_duration', groups=['p'], nan2zero=True)
  4328. # Calculate duration in other zones
  4329. total_social_in_other_duration_df = total_social_duration_df[['B_id']].copy()
  4330. for suffix in ['_pD', '_pS']:
  4331. total_col = f'total_social_duration{suffix}'
  4332. in_far_col = f'in_far_duration{suffix}'
  4333. in_other_col = f'in_other_duration{suffix}'
  4334. if total_col in total_social_duration_df.

analyze_data-checkpoint.ipynb at commit 2b40a15, under Apache-2.0 · at the source

Overview

Authors: Ling-yun Li1, Xinjian Gao2,3, Jun Zhang1,2,3, Wen-wei Wu1,2,3, Ya-tang Li2,3
  1. Department of Neurobiology, School of Basic Medical Sciences, Capital Medical University Beijing China
  2. Beijing Institute for Brain Research, Chinese Academy of Medical Sciences and Peking Union Medical College Beijing China
  3. Chinese Institute for Brain Research Beijing China
Journal: eLife, volume 15, article RP109571
Dates: published online 15 September 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.7554/elife.109571 · PMID 42742131 · PMCID PMC13577663 · OpenAlex W7127583585
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: mouse (organism), rat (organism)
Methods: Statistics
Keywords: innate fear, defensive behavior, social modulation, dominance hierarchy, rodents, Mouse
MeSH: Behavior, Animal*, Fear*, Social Behavior*, Social Dominance*, Animals, Male, Mice, Mice, Inbred C57BL, Rats (* major topic)
Journal subjects: Neuroscience
Topic: Neuroendocrine regulation and behavior (Social Psychology, Psychology), according to OpenAlex
Funding: Natural Science Foundation of Beijing Municipality (5244028, IS23073); National Natural Science Foundation of China (32471071, 32271060); R&D Program of Beijing Municipal Education Commission (1240030201)
Citations: not cited yet (Europe PMC); 67 references in the paper
Research resources: C57BL/6J RRID:IMSR_JAX:000664, MATLAB 2024b RRID:SCR_001622, R v4.2.1 RRID:SCR_001905, Python v3.9 RRID:SCR_008394

Abstract

Fear and defense are among the most fundamental survival behaviors and are profoundly influenced by the social environment in group-living animals. However, it remains poorly understood how social context—and particularly dominance hierarchy, a defining feature of many social species—modulates defensive strategies under naturalistic conditions. To address this question, we investigated the social modulation of innate fear in mice exposed to two ethologically relevant threats: a transient visual looming stimulus and a sustained predatory threat posed by a live rat. We found that social presence alleviated threat-induced stress and modulated defensive behaviors in a rank- and threat-specific manner. During looming exposure, it reduced immediate defensive responses and alleviated post-looming anxiety, with dominants deriving greater benefit. During rat exposure, it promoted a shift from passive to active defense, again most prominently in dominants. These behavioral changes were accompanied by reorganization of transitions between defensive states, indicating that dominance hierarchy shapes both the expression and temporal organization of innate defensive behaviors. Conversely, threat exposure strengthened social engagement, with dominant mice exhibiting more proactive social behaviors and subordinate mice responding more readily to dominant social initiations. Together, these findings demonstrate how dominance hierarchy modulates defensive responses to distinct naturalistic threats and, in turn, how threat experience shapes social behavior, providing a behavioral framework for probing the neural basis of socially modulated innate fear.

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

Repository

Its files are read in the Code ↔ Paper reader above.

YatangLiLab/Gao-2026-rank-threat-soical-modulation

License: Apache-2.0
State: the link answers, verified on 26 September 2026
Evidence: files inventoried
Commit: 2b40a1596de1298b089c74093646b7e29e2671b7, 13 July 2026
Languages: Jupyter (4), Python (4)
Size: 173 files, 8 scripts
Software Heritage: not archived
Found in: “Data availability”
Holds: README, license file, 2 notebooks
Not found: CITATION.cff, environment file, tests, continuous integration, documentation
Tools: Matplotlib (3 files), NumPy (3 files), pandas (3 files), seaborn (3 files), SciPy (2 files)
Availability: 1 check, the latest on 26 September 2026: the link answers
  • 26 September 2026: the link answers
3 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:

  • 1 repository of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
  • 3 scripts, each with its path and the digest of its content;
  • no match between paragraphs and code yet;
  • 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

No dataset and no data link were found in the paper.

Data availability

Data and code are available in a public GitHub repository (https://github.com/YatangLiLab/Gao-2026-rank-threat-soical-modulation; copy archived at Li and Gao, 2026).

Reproduced under the paper's license (CC BY), 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, pages, dates, 5 authors, 6 keywords, 9 MeSH terms, 3 funders, 65 references, 4 RRIDs.

Cite

This paper

Li, L.-y., Gao, X., Zhang, J., Wu, W.-w., & Li, Y.-t. (2026). Rank- and threat-dependent social modulation of innate defensive behaviors. eLife, 15, RP109571. https://doi.org/10.7554/elife.109571

BibTeX

@article{li2026rank,
author = {Li, Ling-yun and Gao, Xinjian and Zhang, Jun and Wu, Wen-wei and Li, Ya-tang},
title = {{Rank- and threat-dependent social modulation of innate defensive behaviors}},
journal = {eLife},
year = {2026},
month = sep,
volume = {15},
pages = {RP109571},
publisher = {eLife Sciences Publications, Ltd},
issn = {2050-084X},
doi = {10.7554/elife.109571},
url = {https://doi.org/10.7554/elife.109571},
pmid = {42742131},
pmcid = {PMC13577663}
}

RIS

TY - JOUR
AU - Li, Ling-yun
AU - Gao, Xinjian
AU - Zhang, Jun
AU - Wu, Wen-wei
AU - Li, Ya-tang
TI - Rank- and threat-dependent social modulation of innate defensive behaviors
T2 - eLife
J2 - Elife
PY - 2026
DA - 2026/09/15
VL - 15
SP - RP109571
SN - 2050-084X
PB - eLife Sciences Publications, Ltd
DO - 10.7554/elife.109571
UR - https://doi.org/10.7554/elife.109571
LA - en
ER -

CSL-JSON

{
"id": "10.7554/elife.109571",
"type": "article-journal",
"title": "Rank- and threat-dependent social modulation of innate defensive behaviors",
"container-title": "eLife",
"author": [
{
"family": "Li",
"given": "Ling-yun"
},
{
"family": "Gao",
"given": "Xinjian"
},
{
"family": "Zhang",
"given": "Jun"
},
{
"family": "Wu",
"given": "Wen-wei"
},
{
"family": "Li",
"given": "Ya-tang"
}
],
"container-title-short": "Elife",
"volume": "15",
"page": "RP109571",
"DOI": "10.7554/elife.109571",
"PMID": "42742131",
"PMCID": "PMC13577663",
"ISSN": "2050-084X",
"publisher": "eLife Sciences Publications, Ltd",
"URL": "https://doi.org/10.7554/elife.109571",
"language": "en",
"issued": {
"date-parts": [
[
2026,
9,
15
]
]
}
}

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

Similar papers

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

[1] doi:10.7554/elife.105528 [code]
Functional specialization of mPFC-BLA and mPFC-NAc pathways in affective state representation.
Journal: eLife
In common: pandas, SciPy, Matplotlib, 1 other tool, mouse, 3 references
[2] doi:10.1016/j.isci.2026.116498 [code]
The Tower Foraging Park: A paradigm for studying cognitive and motor processes underlying behavioral flexibility in freely moving mice.
Journal: iScience
In common: SciPy, Matplotlib, NumPy, mouse, 3 references
[3] doi:10.7554/elife.109717 [code]
Retrosplenial cortex enables context-dependent goal-directed sensorimotor transformation.
Journal: eLife
In common: seaborn, pandas, SciPy, 2 other tools, mouse, 2 references
[4] doi:10.1038/s41467-026-75837-5 [code]
A hierarchical framework for cortical and subcortical gray-matter parcellation across rodents, primates, and humans.
Journal: Nature communications
In common: seaborn, pandas, SciPy, 2 other tools, rat, mouse, 1 reference
[5] doi:10.1038/s41386-026-02492-1 [code]
Proximity in mice induced by an auditory-conditioned stimulus.
Journal: Neuropsychopharmacology : official publication of the American College of Neuropsychopharmacology
In common: seaborn, pandas, SciPy, 2 other tools, mouse, 1 reference
[6] doi:10.1038/s42003-026-10089-z [code]
Medial prefrontal cortex neurons integrate amygdala and hypothalamic oxytocin signals to mediate stress-induced social alterations.
Journal: Communications biology
In common: seaborn, pandas, SciPy, 2 other tools, mouse, 1 reference
[7] doi:10.1016/j.isci.2026.116119 [code]
Distinctly structured social behavior across three rodent strains is associated with different neural activity patterns.
Journal: iScience
In common: seaborn, pandas, SciPy, 2 other tools, rat, mouse, 1 reference
[8] doi:10.1038/s41422-026-01256-2 [code]
Neurovascular coupling in the basolateral amygdala modulates negative emotions.
Journal: Cell research
In common: mouse, 3 references
[9] doi:10.1162/imag.a.1310 [code]
Experimental quality control induces changes in Allen mouse brain connectomes.
Journal: Imaging neuroscience (Cambridge, Mass.)
In common: pandas, SciPy, NumPy, mouse, 2 references
[10] doi:10.1038/s41593-026-02376-z [code]
A framework for comparative analysis of human and mouse cortical neuron dendrites in corresponding brain regions.
Journal: Nature neuroscience
In common: seaborn, pandas, SciPy, 2 other tools, mouse, 1 reference

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.