Rank- and threat-dependent social modulation of innate defensive behaviors.
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
- # %%
- # Reload modules automatically.
- %load_ext autoreload
- %autoreload 2
- import os
- print(os.getcwd())
- import pandas as pd
- import matplotlib.pyplot as plt
- import seaborn as sns
- import numpy as np
- # %% [markdown]
- # ## Load manual labeled data and animal tracking data.
- # %% [markdown]
- # ### Load lst coordinate data and behavior data.
- # %%
- import os
- import pickle
- from analyze_data_utils import lst_read_deepethogram_csv, lst_read_bento_annot, lst_read_decision_csv, lst_read_deeplabcut_h5_and_analyze_data
- # lst_decision_file = os.path.abspath(os.path.join('.', 'lst', 'looming_behavior_sb_1looming_lessn.csv'))
- # lst_deci_dict = lst_read_decision_csv(lst_decision_file)
- # with open('lst_deci_dict.pkl', 'wb') as f:
- # pickle.dump(lst_deci_dict, f)
- with open('lst_deci_dict.pkl', 'rb') as f:
- lst_deci_dict = pickle.load(f)
- # lst_labels_dir = os.path.abspath(os.path.join('.', 'lst', 'deg_csv_data'))
- # lst_bento_corr_file = os.path.abspath(os.path.join('.', 'lst', 'bento_correction_sb_1looming.csv'))
- # lst_bhvr_frames_dict, lst_baseline_bhvr_frames_dict = lst_read_deepethogram_csv(lst_labels_dir, lst_bento_corr_file)
- # with open('lst_bhvr_frames_dict.pkl', 'wb') as f:
- # pickle.dump(lst_bhvr_frames_dict, f)
- # with open('lst_baseline_bhvr_frames_dict.pkl', 'wb') as f:
- # pickle.dump(lst_baseline_bhvr_frames_dict, f)
- with open('lst_bhvr_frames_dict.pkl', 'rb') as f:
- lst_bhvr_frames_dict = pickle.load(f)
- with open('lst_baseline_bhvr_frames_dict.pkl', 'rb') as f:
- lst_baseline_bhvr_frames_dict = pickle.load(f)
- # lst_bento_dir = os.path.abspath(os.path.join('.', 'lst', 'bento_annot_data'))
- # lst_bhvr_labels_dict, lst_baseline_bhvr_labels_dict = lst_read_bento_annot(lst_bento_dir, lst_bento_corr_file, lst_deci_dict)
- # with open('lst_bhvr_labels_dict.pkl', 'wb') as f:
- # pickle.dump(lst_bhvr_labels_dict, f)
- # with open('lst_baseline_bhvr_labels_dict.pkl', 'wb') as f:
- # pickle.dump(lst_baseline_bhvr_labels_dict, f)
- with open('lst_bhvr_labels_dict.pkl', 'rb') as f:
- lst_bhvr_labels_dict = pickle.load(f)
- with open('lst_baseline_bhvr_labels_dict.pkl', 'rb') as f:
- lst_baseline_bhvr_labels_dict = pickle.load(f)
- # lst_dlc_h5_dir = os.path.abspath(os.path.join('.', 'lst', 'dlc_h5_data'))
- # lsts_data, lstp_data = lst_read_deeplabcut_h5_and_analyze_data(lst_dlc_h5_dir)
- # with open('lsts_data.pkl', 'wb') as f:
- # pickle.dump(lsts_data, f)
- # with open('lstp_data.pkl', 'wb') as f:
- # pickle.dump(lstp_data, f)
- with open('lsts_data.pkl', 'rb') as f:
- lsts_data = pickle.load(f)
- with open('lstp_data.pkl', 'rb') as f:
- lstp_data = pickle.load(f)
- # %% [markdown]
- # ### Global variables: ms_ids, mp_ids, looming_start_frames
- # %%
- ms_ids = ['B1M1(S)', 'B1M2(D)', 'B2M1(D)', 'B2M2(S)', 'B3M1(S)', 'B3M2(D)', 'B4M1(D)', 'B4M2(S)',
- 'B5M1(D)', 'B5M2(S)', 'B6M1(D)', 'B6M2(S)', 'B7M1(S)', 'B7M2(D)', 'B8M1(D)', 'B8M2(S)',
- 'B9M1(D)', 'B9M2(S)', 'B10M1(S)', 'B10M2(D)', 'B11M1(S)', 'B11M2(D)', 'B12M1(S)', 'B12M2(D)',
- 'B13M1(S)', 'B13M2(D)', 'B14M1(S)', 'B14M2(D)', 'B15M1(S)', 'B15M2(D)', 'B16M1(S)', 'B16M2(D)',
- 'B17M1(S)', 'B17M2(D)', 'B18M1(D)', 'B18M2(S)', 'B19M1(D)', 'B19M2(S)', 'B20M1(S)', 'B20M2(D)']
- mp_ids = ['B1M1(S)_B1M2(D)', 'B2M1(D)_B2M2(S)', 'B3M1(S)_B3M2(D)', 'B4M1(D)_B4M2(S)',
- 'B5M1(D)_B5M2(S)', 'B6M1(D)_B6M2(S)', 'B7M1(S)_B7M2(D)', 'B8M1(D)_B8M2(S)',
- 'B9M1(D)_B9M2(S)', 'B10M1(S)_B10M2(D)', 'B11M1(S)_B11M2(D)', 'B12M1(S)_B12M2(D)',
- 'B13M1(S)_B13M2(D)', 'B14M1(S)_B14M2(D)', 'B15M1(S)_B15M2(D)', 'B16M1(S)_B16M2(D)',
- 'B17M1(S)_B17M2(D)', 'B18M1(D)_B18M2(S)', 'B19M1(D)_B19M2(S)', 'B20M1(S)_B20M2(D)']
- mS_ids = ['B1M1(S)', 'B2M2(S)', 'B3M1(S)', 'B4M2(S)', 'B5M2(S)', 'B6M2(S)', 'B7M1(S)', 'B8M2(S)',
- 'B9M2(S)', 'B10M1(S)', 'B11M1(S)', 'B12M1(S)', 'B13M1(S)', 'B14M1(S)', 'B15M1(S)', 'B16M1(S)',
- 'B17M1(S)', 'B18M2(S)', 'B19M2(S)', 'B20M1(S)']
- mD_ids = ['B1M2(D)', 'B2M1(D)', 'B3M2(D)', 'B4M1(D)', 'B5M1(D)', 'B6M1(D)', 'B7M2(D)', 'B8M1(D)',
- 'B9M1(D)', 'B10M2(D)', 'B11M2(D)', 'B12M2(D)', 'B13M2(D)', 'B14M2(D)', 'B15M2(D)', 'B16M2(D)',
- 'B17M2(D)', 'B18M1(D)', 'B19M1(D)', 'B20M2(D)']
- looming_start_frames = {'B1M1(S)': [21687, 24882],
- 'B1M2(D)': [19787, 27340],
- 'B2M1(D)': [20822, 28112],
- 'B2M2(S)': [20517, 27347],
- 'B3M1(S)': [23645, 27281],
- 'B3M2(D)': [18638, 26129],
- 'B4M1(D)': [19250, 25754],
- 'B4M2(S)': [19919, 25135],
- 'B5M1(D)': [18124, 25424],
- 'B5M2(S)': [18522, 33582],
- 'B6M1(D)': [18365, 25547, 29965],
- 'B6M2(S)': [18517, 22786, 27173, 34537],
- 'B7M1(S)': [20516, 31677, 45430, 55860, 67593],
- 'B7M2(D)': [18519, 30751, 35634, 40217, 43675],
- 'B8M1(D)': [19927, 22999, 32861, 39477, 47451],
- 'B8M2(S)': [20228, 26561, 34729, 36654, 45224],
- 'B9M1(D)': [19173, 21997, 25862, 28980, 32281],
- 'B9M2(S)': [20031, 24378, 28080, 32859, 37854],
- 'B10M1(S)': [18291, 22143, 25473, 31536, 43488],
- 'B10M2(D)': [19743, 24225, 29010, 34332, 38145],
- 'B11M1(S)': [18180, 22365, 26364, 29664, 32778],
- 'B11M2(D)': [22305, 29619, 35238, 42942, 52587],
- 'B12M1(S)': [26184, 35190, 38490, 44826, 63663],
- 'B12M2(D)': [22101, 26424, 28731, 33624, 38262],
- 'B13M1(S)': [19734, 21690, 30522, 33570, 39768],
- 'B13M2(D)': [19563, 22602, 36426, 43215, 49560],
- 'B14M1(S)': [21639, 24939, 29826, 33582, 38907],
- 'B14M2(D)': [25443, 28803, 32490, 42564, 52275],
- 'B15M1(S)': [18798, 24028, 27337, 32670, 39426, 45213, 50346, 53076],
- 'B15M2(D)': [18759, 25234, 30829, 34297, 37946, 43971, 48324, 51170, 64290, 69035],
- 'B16M1(S)': [18354, 20851, 25998, 31204, 35160, 40137, 51816],
- 'B16M2(D)': [18282, 25614, 28758, 33453, 35790, 41724, 52158, 69162],
- 'B17M1(S)': [18922, 22083, 25744, 33987, 37318, 44458, 53217, 61588, 65379],
- 'B17M2(D)': [18907, 23739, 30816, 41705, 53482],
- 'B18M1(D)': [18645, 21592, 45597, 52743],
- 'B18M2(S)': [18899, 21992, 28926, 36918, 40850, 54897, 58513, 61513],
- 'B19M1(D)': [19371, 27168, 42195, 50253, 58077],
- 'B19M2(S)': [19974, 28379],
- 'B20M1(S)': [22228, 26274, 55110, 60789, 71464],
- 'B20M2(D)': [18991, 24878, 28582, 35260, 38830, 45945, 51220, 55177, 60471, 64398],
- 'B1M1(S)_B1M2(D)': [26897, 32648],
- 'B2M1(D)_B2M2(S)': [53137, 60549, 70481],
- 'B3M1(S)_B3M2(D)': [29541, 37526, 40262, 47454, 69994],
- 'B5M1(D)_B5M2(S)': [25255, 28961, 39435, 43291, 70922],
- 'B4M1(D)_B4M2(S)': [20228, 23426, 26259, 28705, 33779],
- 'B6M1(D)_B6M2(S)': [25041, 29475, 40256, 45849, 63113],
- 'B7M1(S)_B7M2(D)': [22068, 33550, 52853, 57657, 65416],
- 'B8M1(D)_B8M2(S)': [22155, 31191, 34569, 51364, 70682],
- 'B9M1(D)_B9M2(S)': [18831, 20865, 24561, 28755, 33870, 45141, 49365, 52656, 69148],
- 'B10M1(S)_B10M2(D)': [22422, 27787, 36190, 41217, 45228, 48732, 52455, 56308, 60924, 66468],
- 'B11M1(S)_B11M2(D)': [19539, 23068, 26653, 30262, 34231, 37920, 41037, 46828, 49966, 53490],
- 'B12M1(S)_B12M2(D)': [19257, 23934, 28201, 30547, 34242, 37833, 42060, 47703, 52767, 57060],
- 'B13M1(S)_B13M2(D)': [19183, 23344, 29673, 33379, 35610, 41544, 48450, 51997, 56022, 63910],
- 'B14M1(S)_B14M2(D)': [22887, 29199, 40301, 48266, 52510, 67180, 71367],
- 'B15M1(S)_B15M2(D)': [18831, 22671, 28074, 31942, 36654, 39847, 44067, 49578, 53272, 57154],
- 'B16M1(S)_B16M2(D)': [20454, 29139, 32550, 36150, 40696, 44059, 49290, 52071, 58332, 63426],
- 'B17M1(S)_B17M2(D)': [25080, 33249, 37938, 45118, 51993, 55266, 60331, 63778, 68193, 71500],
- 'B18M1(D)_B18M2(S)': [18876, 24573, 28573, 32257, 36224, 38827, 42231, 44521, 49630, 56858],
- 'B19M1(D)_B19M2(S)': [28854, 33918, 37251, 41328, 46119, 49065, 51954, 56518, 61744, 64632],
- 'B20M1(S)_B20M2(D)': [22090, 27171, 31389, 37048, 40728, 45769, 48780, 52728, 54745, 58203]}
- lst_pixpercm = 21.76
- framerate = 30
- # %%
- lst_all_bhvrs = [
- 'escape',
- 'freezing',
- 'tail_rattling',
- 'rearing/up_stretch',
- 'huddling',
- 'approach_partner',
- 'follow_partner',
- 'groom_partner',
- 'sniff_partner',
- 'grooming',
- 'sniffing',
- 'gap_state',
- ]
- lst_bhvr_abbrev_dict = {
- "escape": 'E',
- "freezing": 'F',
- "tail_rattling": 'T',
- "rearing/up_stretch": 'R',
- "huddling": 'H',
- "approach_partner": 'AP',
- "follow_partner": 'FP',
- "groom_partner": 'GP',
- "sniff_partner": 'SP',
- "grooming": 'G',
- "sniffing": 'S',
- 'gap_state': 'O'
- }
- lst_bhvr_color_dict = {
- "E": "#FF1717",
- "R": "#C409FF",
- "F": "#FF75EF",
- "T": "#FF09FF",
- "H": "#6D57F3",
- "AP": "#143FCA",
- "FP": "#137CAB",
- "GP": "#0990FF",
- "SP": "#09FFFF",
- "G": "#2EA761",
- "S": "#0FD400",
- "O": "white",
- 'Def': 'red',
- 'Soc': 'royalblue',
- 'Exp&\nRest': 'limegreen'
- }
- lst_arena_params={
- 'pixpergrid': 43.52,
- 'arena_size': 1088,
- 'edge_width': 43.52 * 4,
- 'nest_coord': [0, 43.52 * 17.5],
- 'nest_size': [43.52 * 10, 43.52 * 7.5]
- }
- # %% [markdown]
- # ### Move csv files and convert csv to annot.
- # %%
- import os
- from move_csv_files import copy_csv_files_to_root, move_csv_files_to_target_folders
- lsts_deepethogram_project_DATA_dir = os.path.abspath(os.path.join('.', 'lst', 'lsts_deepethogram', 'DATA'))
- copy_csv_files_to_root(lsts_deepethogram_project_DATA_dir)
- lstp_deepethogram_project_DATA_dir = os.path.abspath(os.path.join('.', 'lst', 'lstp_deepethogram', 'DATA'))
- copy_csv_files_to_root(lstp_deepethogram_project_DATA_dir)
- lst_labels_dir = os.path.abspath(os.path.join('.', 'lst', 'deg_csv_data'))
- temp_lsts_labels_dir = os.path.abspath(os.path.join('.', 'lst', 'labels_csv_lsts'))
- temp_lstp_labels_dir = os.path.abspath(os.path.join('.', 'lst', 'labels_csv_lstp'))
- move_csv_files_to_target_folders(lsts_deepethogram_project_DATA_dir, temp_lsts_labels_dir)
- move_csv_files_to_target_folders(lstp_deepethogram_project_DATA_dir, temp_lstp_labels_dir)
- # %%
- from convert_labelscsv2annot_lst import get_csv_files, read_labelscsv_file, create_annot_content, creat_annot_file
- lst_bento_dir = os.path.abspath(os.path.join('.', 'lst', 'bento_annot_data'))
- group = 's'
- csv_files = get_csv_files(temp_lsts_labels_dir)
- for csv_file in csv_files:
- csv_file_path = os.path.join(temp_lsts_labels_dir, csv_file)
- behavior_frames_dict = read_labelscsv_file(csv_file_path, group)
- annot_content = create_annot_content(csv_file, group)
- creat_annot_file(csv_file, lst_bento_dir, annot_content, behavior_frames_dict)
- group = 'p'
- csv_files = get_csv_files(temp_lstp_labels_dir)
- for csv_file in csv_files:
- csv_file_path = os.path.join(temp_lstp_labels_dir, csv_file)
- behavior_frames_dict = read_labelscsv_file(csv_file_path, group)
- annot_content = create_annot_content(csv_file, group)
- creat_annot_file(csv_file, lst_bento_dir, annot_content, behavior_frames_dict)
- # %%
- move_csv_files_to_target_folders(temp_lsts_labels_dir, lst_labels_dir)
- move_csv_files_to_target_folders(temp_lstp_labels_dir, lst_labels_dir)
- # %% [markdown]
- # ### Check lst overlap.
- # %%
- from collections import Counter
- # Set frame range
- frame_start = 150
- frame_end = 1950
- total_frames = frame_end - frame_start
- # Excluded behaviors
- excluded_behaviors = ['nest']
- # Iterate through all groups and ranks
- for group in ['s', 'p']:
- group_name = 'Single' if group == 's' else 'Pair'
- print(f"\n{'='*60}")
- print(f"Group: {group_name}")
- print('='*60)
- for rank in ['D', 'S']:
- rank_name = 'Dominant' if rank == 'D' else 'Subordinate'
- print(f"\n--- Rank: {rank_name} ---")
- if group not in lst_bhvr_frames_dict or rank not in lst_bhvr_frames_dict[group]:
- print(f" No data for {group_name} {rank_name}")
- continue
- # Initialize aggregated counter for all trials in this group-rank combination
- aggregated_combinations = []
- aggregated_overlap_frames = [] # Store all overlapping frames across trials
- trial_count = 0
- # Iterate through all trials
- for t_id in sorted(lst_bhvr_frames_dict[group][rank].keys()):
- trial_data = lst_bhvr_frames_dict[group][rank][t_id]
- # Get all available behaviors for this trial
- available_behaviors = [b for b in trial_data.keys() if b not in excluded_behaviors]
- if len(available_behaviors) == 0:
- print(f"\n {t_id}: No behaviors available after exclusion")
- continue
- # Create a list to store behavior combination for each frame
- frame_combinations = []
- overlap_frames_info = [] # Store frame numbers with overlapping behaviors
- # Iterate through each frame
- for frame_idx in range(frame_start, frame_end):
- # Find all active behaviors at this frame
- active_behaviors = []
- for behavior in sorted(available_behaviors): # Sort for consistent ordering
- try:
- if frame_idx in trial_data[behavior].index and trial_data[behavior][frame_idx] == 1:
- active_behaviors.append(behavior)
- except:
- pass
- # priority_behaviors = ['stretching']
- # if 'sniffing' in active_behaviors:
- # for priority_bhvr in priority_behaviors:
- # if priority_bhvr in active_behaviors:
- # active_behaviors.remove('sniffing')
- # break
- # Create combination string
- if len(active_behaviors) == 0:
- combo_str = "no_behavior"
- elif len(active_behaviors) == 1:
- combo_str = active_behaviors[0]
- else:
- combo_str = "+".join(active_behaviors)
- # Record overlapping frame information
- overlap_frames_info.append({
- 'frame': frame_idx,
- 'behaviors': active_behaviors
- })
- frame_combinations.append(combo_str)
- # Add this trial's combinations to the aggregated list
- aggregated_combinations.extend(frame_combinations)
- # Add overlapping frames info with trial ID
- for frame_info in overlap_frames_info:
- aggregated_overlap_frames.append({
- 'trial': t_id,
- 'frame': frame_info['frame'],
- 'behaviors': frame_info['behaviors']
- })
- trial_count += 1
- # Count occurrences of each combination
- combination_counter = Counter(frame_combinations)
- # Calculate percentages and sort by frequency
- combination_stats = []
- for combo, count in combination_counter.items():
- pct = (count / total_frames) * 100
- combination_stats.append({
- 'combination': combo,
- 'frames': count,
- 'pct': pct
- })
- # Sort by frame count (descending)
- combination_stats.sort(key=lambda x: x['frames'], reverse=True)
- # Print results
- print(f"\n {t_id}:")
- print(f" Total unique behavior combinations: {len(combination_stats)}")
- print(f" Frame distribution:")
- # Show all combinations
- for stat in combination_stats:
- print(f" {stat['combination']}: {stat['frames']} frames ({stat['pct']:.2f}%)")
- # Display overlapping frames details
- if len(overlap_frames_info) > 0:
- print(f"\n Overlapping frames details ({len(overlap_frames_info)} frames total):")
- # Show first 20 overlapping frames as examples
- max_display = min(20, len(overlap_frames_info))
- for i, frame_info in enumerate(overlap_frames_info[:max_display]):
- behaviors_str = " + ".join(frame_info['behaviors'])
- print(f" Frame {frame_info['frame']}: {behaviors_str}")
- if len(overlap_frames_info) > max_display:
- print(f" ... and {len(overlap_frames_info) - max_display} more overlapping frames")
- # Summary: single vs multiple behaviors
- single_behavior_frames = sum(stat['frames'] for stat in combination_stats
- if '+' not in stat['combination'] and stat['combination'] != 'no_behavior')
- multiple_behavior_frames = sum(stat['frames'] for stat in combination_stats
- if '+' in stat['combination'])
- no_behavior_frames = sum(stat['frames'] for stat in combination_stats
- if stat['combination'] == 'no_behavior')
- print(f"\n Summary:")
- print(f" Single behavior: {single_behavior_frames} frames ({single_behavior_frames/total_frames*100:.2f}%)")
- print(f" Multiple behaviors: {multiple_behavior_frames} frames ({multiple_behavior_frames/total_frames*100:.2f}%)")
- print(f" No behavior: {no_behavior_frames} frames ({no_behavior_frames/total_frames*100:.2f}%)")
- # Print aggregated statistics for this group-rank combination
- if trial_count > 0:
- print(f"\n{'*' * 50}")
- print(f"AGGREGATED STATISTICS: {group_name} - {rank_name}")
- print(f"Total trials: {trial_count}")
- print(f"Total frames analyzed: {len(aggregated_combinations)}")
- print(f"{'*' * 50}")
- # Count aggregated combinations
- aggregated_counter = Counter(aggregated_combinations)
- total_aggregated_frames = len(aggregated_combinations)
- # Calculate percentages and sort
- aggregated_stats = []
- for combo, count in aggregated_counter.items():
- pct = (count / total_aggregated_frames) * 100
- aggregated_stats.append({
- 'combination': combo,
- 'frames': count,
- 'pct': pct
- })
- aggregated_stats.sort(key=lambda x: x['frames'], reverse=True)
- print(f"\nAggregated behavior combinations:")
- for stat in aggregated_stats:
- print(f" {stat['combination']}: {stat['frames']} frames ({stat['pct']:.2f}%)")
- # Display all overlapping frames in aggregated data
- if len(aggregated_overlap_frames) > 0:
- print(f"\n All overlapping frames across trials ({len(aggregated_overlap_frames)} frames total):")
- # Group by trial for better readability
- current_trial = None
- max_display_per_trial = 10
- trial_frame_count = 0
- for frame_info in aggregated_overlap_frames:
- if current_trial != frame_info['trial']:
- if current_trial is not None and trial_frame_count > max_display_per_trial:
- print(f" ... and {trial_frame_count - max_display_per_trial} more frames in this trial")
- current_trial = frame_info['trial']
- trial_frame_count = 0
- print(f"\n {current_trial}:")
- trial_frame_count += 1
- if trial_frame_count <= max_display_per_trial:
- behaviors_str = " + ".join(frame_info['behaviors'])
- print(f" Frame {frame_info['frame']}: {behaviors_str}")
- if trial_frame_count > max_display_per_trial:
- print(f" ... and {trial_frame_count - max_display_per_trial} more frames in this trial")
- # Aggregated summary
- agg_single = sum(stat['frames'] for stat in aggregated_stats
- if '+' not in stat['combination'] and stat['combination'] != 'no_behavior')
- agg_multiple = sum(stat['frames'] for stat in aggregated_stats
- if '+' in stat['combination'])
- agg_no_behavior = sum(stat['frames'] for stat in aggregated_stats
- if stat['combination'] == 'no_behavior')
- print(f"\nAggregated Summary:")
- print(f" Single behavior: {agg_single} frames ({agg_single/total_aggregated_frames*100:.2f}%)")
- print(f" Multiple behaviors: {agg_multiple} frames ({agg_multiple/total_aggregated_frames*100:.2f}%)")
- print(f" No behavior: {agg_no_behavior} frames ({agg_no_behavior/total_aggregated_frames*100:.2f}%)")
- print(f"{'*' * 50}\n")
- print("\n" + "="*80)
- print("Frame-by-frame analysis complete")
- print("="*80)
- # %% [markdown]
- # ### Load ret coordinate data and behavior data.
- # %%
- import os
- import pickle
- from analyze_data_utils import ret_read_deepethogram_csv, ret_read_bento_annot, ret_read_ethovision_xlsx_and_analyze_data, reorder_dict_sessions
- session_order = ['B1', 'B2', 'B3', 'B4', 'B5', 'B6', 'B7', 'B8', 'B9', 'B10', 'B11']
- # ret_labels_dir = os.path.abspath(os.path.join('.', 'ret', 'deg_csv_data'))
- # ret_bento_corr_file = os.path.abspath(os.path.join('.', 'ret', 'bento_correction_sb_ret.csv'))
- # ret_bhvr_frames_dict, ret_origin_bhvr_frames_dict = ret_read_deepethogram_csv(ret_labels_dir, ret_bento_corr_file)
- # ret_bhvr_frames_dict = reorder_dict_sessions(ret_bhvr_frames_dict, session_order)
- # ret_origin_bhvr_frames_dict = reorder_dict_sessions(ret_origin_bhvr_frames_dict, session_order)
- # with open('ret_bhvr_frames_dict.pkl', 'wb') as f:
- # pickle.dump(ret_bhvr_frames_dict, f)
- # with open('ret_origin_bhvr_frames_dict.pkl', 'wb') as f:
- # pickle.dump(ret_origin_bhvr_frames_dict, f)
- with open('ret_bhvr_frames_dict.pkl', 'rb') as f:
- ret_bhvr_frames_dict = pickle.load(f)
- with open('ret_origin_bhvr_frames_dict.pkl', 'rb') as f:
- ret_origin_bhvr_frames_dict = pickle.load(f)
- # ret_bento_dir = os.path.abspath(os.path.join('.', 'ret', 'bento_annot_data'))
- # ret_bhvr_labels_dict, ret_origin_bhvr_labels_dict = ret_read_bento_annot(ret_bento_dir, ret_bento_corr_file)
- # ret_bhvr_labels_dict = reorder_dict_sessions(ret_bhvr_labels_dict, session_order)
- # ret_origin_bhvr_labels_dict = reorder_dict_sessions(ret_origin_bhvr_labels_dict, session_order)
- # with open('ret_bhvr_labels_dict.pkl', 'wb') as f:
- # pickle.dump(ret_bhvr_labels_dict, f)
- # with open('ret_origin_bhvr_labels_dict.pkl', 'wb') as f:
- # pickle.dump(ret_origin_bhvr_labels_dict, f)
- with open('ret_bhvr_labels_dict.pkl', 'rb') as f:
- ret_bhvr_labels_dict = pickle.load(f)
- with open('ret_origin_bhvr_labels_dict.pkl', 'rb') as f:
- ret_origin_bhvr_labels_dict = pickle.load(f)
- # ret_ev_xlsx_dir = os.path.abspath(os.path.join('.', 'ret', 'ev_xlsx_data'))
- # ret_raw_data, ret_stat_data = ret_read_ethovision_xlsx_and_analyze_data(ret_ev_xlsx_dir)
- # with open('ret_raw_data.pkl', 'wb') as f:
- # pickle.dump(ret_raw_data, f)
- # with open('ret_stat_data.pkl', 'wb') as f:
- # pickle.dump(ret_stat_data, f)
- with open('ret_raw_data.pkl', 'rb') as f:
- ret_raw_data = pickle.load(f)
- with open('ret_stat_data.pkl', 'rb') as f:
- ret_stat_data = pickle.load(f)
- # %% [markdown]
- # ### Global variables: rat_in_frames
- # %%
- rat_in_frames = {
- 'D1M1&M2': 11100,
- 'D1M1': 10766,
- 'D1M2': 10713,
- 'D2M1&M2': 10416,
- 'D2M1': 10651,
- 'D2M2': 10633,
- 'D3M1&M2': 10462,
- 'D3M1': 10987,
- 'D3M2': 10677,
- 'D4M1&M2': 12051,
- 'D4M1': 11359,
- 'D4M2': 10424,
- 'D5M1&M2': 11100,
- 'D5M1': 10414,
- 'D5M2': 10395,
- 'D6M1&M2': 10751,
- 'D6M1': 10301,
- 'D6M2': 10420,
- 'D7M1&M2': 10374,
- 'D7M1': 10141,
- 'D7M2': 9586,
- 'D8M1&M2': 10082,
- 'D8M1': 10098,
- 'D8M2': 10197,
- 'D9M1&M2': 10188,
- 'D9M1': 10283,
- 'D9M2': 10441,
- 'D10M1&M2': 10510,
- 'D10M1': 10154,
- 'D10M2': 10343,
- 'D11M1&M2': 10463,
- 'D11M1': 10176,
- 'D11M2': 10413
- }
- framerate = 30
- # %%
- offset_frames = {
- 'D1M1&M2': 10812,
- 'D1M1': 10766,
- 'D1M2': 10713,
- 'D2M1&M2': 10286,
- 'D2M1': 10651,
- 'D2M2': 10633,
- 'D3M1&M2': 10470,
- 'D3M1': 10987,
- 'D3M2': 10677,
- 'D4M1&M2': 11817,
- 'D4M1': 11359,
- 'D4M2': 10424,
- 'D5M1&M2': 10976,
- 'D5M1': 10414,
- 'D5M2': 10395,
- 'D6M1&M2': 10624,
- 'D6M1': 10301,
- 'D6M2': 10420,
- 'D7M1&M2': 10225,
- 'D7M1': 10141,
- 'D7M2': 9586,
- 'D8M1&M2': 10066,
- 'D8M1': 10098,
- 'D8M2': 10197,
- 'D9M1&M2': 10196,
- 'D9M1': 10283,
- 'D9M2': 10441,
- 'D10M1&M2': 10556,
- 'D10M1': 10154,
- 'D10M2': 10343,
- 'D11M1&M2': 10942,
- 'D11M1': 10176,
- 'D11M2': 10413
- }
- # %%
- ret_all_bhvrs = [
- "approach",
- "investigation",
- "withdrawal",
- "stretch-attend",
- "freezing",
- "tail_rattling",
- "huddling",
- "approach_partner",
- "follow_partner",
- "groom_partner",
- "sniff_partner",
- "grooming",
- "gap_state",
- ]
- ret_bhvr_color_dict = {
- "A": "#FFCA09",
- "A+AP": "#8A856A",
- "I": "#FF8409",
- "W": "#FF0990",
- "W+AP": "#8A24AD",
- "S": "#FF7979",
- "F": "#FF75EF",
- "T": "#FF09FF",
- "H": "#6D57F3",
- "H+F": "#B666F1",
- "AP": "#143FCA",
- "FP": "#137CAB",
- "GP": "#0990FF",
- "SP": "#09FFFF",
- "G": "#28AE61",
- "O": "white",
- 'Def': 'red',
- 'Soc': 'royalblue',
- 'Exp&\nRest': 'limegreen'
- }
- ret_bhvr_abbrev_dict = {
- "approach": "A",
- "investigation": "I",
- "withdrawal": "W",
- "stretch-attend": "S",
- "freezing": "F",
- "tail_rattling": "T",
- "huddling": "H",
- "approach_partner": "AP",
- "follow_partner": "FP",
- "groom_partner": "GP",
- "sniff_partner": "SP",
- "grooming": "G",
- 'gap_state': 'O'
- }
- ret_arena_params={
- 'dx': 24/24, # pixel per cm
- 'rat_x': -5,
- 'near_width': 9,
- 'middle_width': 9,
- 'far_width': 6,
- 'arena_size': 24 * (24/24)
- }
- # %% [markdown]
- # ### Move csv files and convert csv to annot.
- # %%
- import os
- from move_csv_files import copy_csv_files_to_root, move_csv_files_to_target_folders, rename_csv_files_remove_video_prefix
- rets_deepethogram_project_DATA_dir = os.path.abspath(os.path.join('.', 'ret', 'rets_deepethogram', 'DATA'))
- copy_csv_files_to_root(rets_deepethogram_project_DATA_dir)
- retp_deepethogram_project_DATA_dir = os.path.abspath(os.path.join('.', 'ret', 'retp_deepethogram', 'DATA'))
- copy_csv_files_to_root(retp_deepethogram_project_DATA_dir)
- ret_labels_dir = os.path.abspath(os.path.join('.', 'ret', 'deg_csv_data'))
- temp_rets_labels_dir = os.path.abspath(os.path.join('.', 'ret', 'labels_csv_rets'))
- temp_retp_labels_dir = os.path.abspath(os.path.join('.', 'ret', 'labels_csv_retp'))
- move_csv_files_to_target_folders(rets_deepethogram_project_DATA_dir, temp_rets_labels_dir)
- move_csv_files_to_target_folders(retp_deepethogram_project_DATA_dir, temp_retp_labels_dir)
- rename_csv_files_remove_video_prefix(temp_rets_labels_dir)
- rename_csv_files_remove_video_prefix(temp_retp_labels_dir)
- # %%
- from convert_labelscsv2annot_ret import get_csv_files, read_labelscsv_file, create_annot_content, creat_annot_file
- ret_bento_dir = os.path.abspath(os.path.join('.', 'ret', 'bento_annot_data'))
- group = 's'
- csv_files = get_csv_files(temp_rets_labels_dir)
- for csv_file in csv_files:
- csv_file_path = os.path.join(temp_rets_labels_dir, csv_file)
- behavior_frames_dict = read_labelscsv_file(csv_file_path)
- annot_content = create_annot_content(csv_file)
- creat_annot_file(csv_file, ret_bento_dir, annot_content, behavior_frames_dict)
- group = 'p'
- csv_files = get_csv_files(temp_retp_labels_dir)
- for csv_file in csv_files:
- csv_file_path = os.path.join(temp_retp_labels_dir, csv_file)
- behavior_frames_dict = read_labelscsv_file(csv_file_path)
- annot_content = create_annot_content(csv_file)
- creat_annot_file(csv_file, ret_bento_dir, annot_content, behavior_frames_dict)
- # %%
- move_csv_files_to_target_folders(temp_rets_labels_dir, ret_labels_dir)
- move_csv_files_to_target_folders(temp_retp_labels_dir, ret_labels_dir)
- # %% [markdown]
- # ### Check ret overlap.
- # %%
- from collections import Counter
- # Control variable: if True, check all frames; if False, use specified frame range
- if_all_frame = False
- # Set frame range
- if if_all_frame:
- # Get the maximum frame number from all trials
- max_frame = 8900
- for group in ['s', 'p']:
- if group in ret_origin_bhvr_frames_dict:
- for rank in ['D', 'S']:
- if rank in ret_origin_bhvr_frames_dict[group]:
- for t_id in ret_origin_bhvr_frames_dict[group][rank].keys():
- trial_data = ret_origin_bhvr_frames_dict[group][rank][t_id]
- for behavior in trial_data.keys():
- if len(trial_data[behavior]) > 0:
- max_frame = max(max_frame, trial_data[behavior].index.max())
- frame_start = 0
- frame_end = max_frame + 1
- print("\n" + "="*80)
- print(f"Frame-by-Frame Behavior Combination Analysis (ALL frames: 0-{frame_end})")
- print("="*80)
- else:
- frame_start = 0
- frame_end = 8900
- print("\n" + "="*80)
- print("Frame-by-Frame Behavior Combination Analysis (0:8900 frames)")
- print("="*80)
- total_frames = frame_end - frame_start
- # Excluded behaviors
- excluded_behaviors = ['rat_in', 'in_proximity', 'stretching/rearing', 'stretching/rearing.1']
- # Iterate through all groups and ranks
- for group in ['s', 'p']:
- group_name = 'Single' if group == 's' else 'Pair'
- print(f"\n{'='*60}")
- print(f"Group: {group_name} (Frames: {frame_start}-{frame_end})")
- print('='*60)
- for rank in ['D', 'S']:
- rank_name = 'Dominant' if rank == 'D' else 'Subordinate'
- print(f"\n--- Rank: {rank_name} ---")
- if group not in ret_origin_bhvr_frames_dict or rank not in ret_origin_bhvr_frames_dict[group]:
- print(f" No data for {group_name} {rank_name}")
- continue
- # Initialize aggregated counter for all trials in this group-rank combination
- aggregated_combinations = []
- aggregated_overlap_frames = [] # Store all overlapping frames across trials
- trial_count = 0
- # Iterate through all trials
- for t_id in sorted(ret_origin_bhvr_frames_dict[group][rank].keys()):
- trial_data = ret_origin_bhvr_frames_dict[group][rank][t_id]
- # Get all available behaviors for this trial
- available_behaviors = [b for b in trial_data.keys() if b not in excluded_behaviors]
- if len(available_behaviors) == 0:
- print(f"\n {t_id}: No behaviors available after exclusion")
- continue
- # Create a list to store behavior combination for each frame
- frame_combinations = []
- overlap_frames_info = [] # Store frame numbers with overlapping behaviors
- # Iterate through each frame
- for frame_idx in range(frame_start, frame_end):
- # Find all active behaviors at this frame
- active_behaviors = []
- for behavior in sorted(available_behaviors): # Sort for consistent ordering
- try:
- if frame_idx in trial_data[behavior].index and trial_data[behavior][frame_idx] == 1:
- active_behaviors.append(behavior)
- except:
- pass
- # Create combination string
- if len(active_behaviors) == 0:
- combo_str = "no_behavior"
- elif len(active_behaviors) == 1:
- combo_str = active_behaviors[0]
- else:
- combo_str = "+".join(active_behaviors)
- # Record overlapping frame information
- overlap_frames_info.append({
- 'frame': frame_idx,
- 'behaviors': active_behaviors
- })
- frame_combinations.append(combo_str)
- # Add this trial's combinations to the aggregated list
- aggregated_combinations.extend(frame_combinations)
- # Add overlapping frames info with trial ID
- for frame_info in overlap_frames_info:
- aggregated_overlap_frames.append({
- 'trial': t_id,
- 'frame': frame_info['frame'],
- 'behaviors': frame_info['behaviors']
- })
- trial_count += 1
- # Count occurrences of each combination
- combination_counter = Counter(frame_combinations)
- # Calculate percentages and sort by frequency
- combination_stats = []
- for combo, count in combination_counter.items():
- pct = (count / total_frames) * 100
- combination_stats.append({
- 'combination': combo,
- 'frames': count,
- 'pct': pct
- })
- # Sort by frame count (descending)
- combination_stats.sort(key=lambda x: x['frames'], reverse=True)
- # Print results
- print(f"\n {t_id}:")
- print(f" Total unique behavior combinations: {len(combination_stats)}")
- print(f" Frame distribution:")
- # Show all combinations
- for stat in combination_stats:
- print(f" {stat['combination']}: {stat['frames']} frames ({stat['pct']:.2f}%)")
- # Display overlapping frames details
- if len(overlap_frames_info) > 0:
- print(f"\n Overlapping frames details ({len(overlap_frames_info)} frames total):")
- for frame_info in overlap_frames_info:
- behaviors_str = " + ".join(frame_info['behaviors'])
- print(f" Frame {frame_info['frame']}: {behaviors_str}")
- # Summary: single vs multiple behaviors
- single_behavior_frames = sum(stat['frames'] for stat in combination_stats
- if '+' not in stat['combination'] and stat['combination'] != 'no_behavior')
- multiple_behavior_frames = sum(stat['frames'] for stat in combination_stats
- if '+' in stat['combination'])
- no_behavior_frames = sum(stat['frames'] for stat in combination_stats
- if stat['combination'] == 'no_behavior')
- print(f"\n Summary:")
- print(f" Single behavior: {single_behavior_frames} frames ({single_behavior_frames/total_frames*100:.2f}%)")
- print(f" Multiple behaviors: {multiple_behavior_frames} frames ({multiple_behavior_frames/total_frames*100:.2f}%)")
- print(f" No behavior: {no_behavior_frames} frames ({no_behavior_frames/total_frames*100:.2f}%)")
- # Print aggregated statistics for this group-rank combination
- if trial_count > 0:
- print(f"\n{'*' * 50}")
- print(f"AGGREGATED STATISTICS: {group_name} - {rank_name}")
- print(f"Total trials: {trial_count}")
- print(f"Total frames analyzed: {len(aggregated_combinations)}")
- print(f"{'*' * 50}")
- # Count aggregated combinations
- aggregated_counter = Counter(aggregated_combinations)
- total_aggregated_frames = len(aggregated_combinations)
- # Calculate percentages and sort
- aggregated_stats = []
- for combo, count in aggregated_counter.items():
- pct = (count / total_aggregated_frames) * 100
- aggregated_stats.append({
- 'combination': combo,
- 'frames': count,
- 'pct': pct
- })
- aggregated_stats.sort(key=lambda x: x['frames'], reverse=True)
- print(f"\nAggregated behavior combinations:")
- for stat in aggregated_stats:
- print(f" {stat['combination']}: {stat['frames']} frames ({stat['pct']:.2f}%)")
- # Display all overlapping frames in aggregated data
- if len(aggregated_overlap_frames) > 0:
- print(f"\n All overlapping frames across trials ({len(aggregated_overlap_frames)} frames total):")
- # Group by trial for better readability
- current_trial = None
- for frame_info in aggregated_overlap_frames:
- if current_trial != frame_info['trial']:
- current_trial = frame_info['trial']
- print(f"\n {current_trial}:")
- behaviors_str = " + ".join(frame_info['behaviors'])
- print(f" Frame {frame_info['frame']}: {behaviors_str}")
- # Aggregated summary
- agg_single = sum(stat['frames'] for stat in aggregated_stats
- if '+' not in stat['combination'] and stat['combination'] != 'no_behavior')
- agg_multiple = sum(stat['frames'] for stat in aggregated_stats
- if '+' in stat['combination'])
- agg_no_behavior = sum(stat['frames'] for stat in aggregated_stats
- if stat['combination'] == 'no_behavior')
- print(f"\nAggregated Summary:")
- print(f" Single behavior: {agg_single} frames ({agg_single/total_aggregated_frames*100:.2f}%)")
- print(f" Multiple behaviors: {agg_multiple} frames ({agg_multiple/total_aggregated_frames*100:.2f}%)")
- print(f" No behavior: {agg_no_behavior} frames ({agg_no_behavior/total_aggregated_frames*100:.2f}%)")
- print(f"{'*' * 50}\n")
- print("\n" + "="*80)
- print("Frame-by-frame analysis complete")
- print("="*80)
- # %% [markdown]
- # ## Analyze data and plot figures.
- # %% [markdown]
- # ### lst: lot ethogram and velocity
- # %%
- # Plot ethogram and velocity for all trials
- import matplotlib.pyplot as plt
- import matplotlib.patches as mpatches
- import os
- # Create output directory if it doesn't exist
- os.makedirs('data/ethogram_velocity_plots', exist_ok=True)
- # Behavior colors
- behavior_colors = {'escape': 'red', 'freezing': 'blue', 'sniffing': 'green',
- 'grooming': 'purple', 'stretching': 'orange', 'rearing': 'brown'}
- # Loop through all trials
- plot_count = 0
- for g in ['s', 'p']:
- for r in ['D', 'S']:
- for t_id in lst_bhvr_frames_dict[g][r].keys():
- try:
- # Parse m_id from t_id
- if t_id.split('B')[1].split('M')[0] in ['1', '3', '7', '10', '11', '12', '13', '14', '15', '16', '17', '20']:
- m = 1 if t_id.split('M')[1].split('T')[0] == 'S' else 2
- else:
- m = 1 if t_id.split('M')[1].split('T')[0] == 'D' else 2
- m_id = t_id.split('M')[0] + 'M' + str(m) + '(' + t_id.split('M')[1].split('T')[0] + ')'
- # Get velocity data and looming_start_frame based on group
- if g == 's':
- velocity_data = lsts_data[m_id]['velocity']['waist']
- looming_start_frame = looming_start_frames[m_id][int(t_id.split('T')[1])-1]
- elif g == 'p':
- velocity_data = lstp_data[m_id]['velocity']['waist']
- mp_id = mp_ids[int(m_id.split('B')[1].split('M')[0])-1]
- looming_start_frame = looming_start_frames[mp_id][int(t_id.split('T')[1])-1]
- # Get behavior data for this trial
- trial_bhvr_data = lst_bhvr_frames_dict[g][r][t_id]
- # Define time window: -5s to +60s (150 frames before to 1800 frames after looming start)
- plot_start_frame = looming_start_frame - 150
- plot_end_frame = looming_start_frame + 1800
- plot_window = slice(plot_start_frame, plot_end_frame)
- # Extract velocity data for plotting (convert to cm/s)
- velocity_segment = velocity_data[plot_window].values
- velocity_cm_s = velocity_segment / lst_pixpercm * framerate
- time_axis = np.arange(len(velocity_cm_s)) / framerate # Convert to seconds
- # Get all behaviors
- behaviors = [b for b in trial_bhvr_data.keys() if b != 'nest']
- # Create figure with two subplots
- fig, (ax1, ax2) = plt.subplots(2, 1, figsize=(16, 8), sharex=True,
- gridspec_kw={'height_ratios': [1, 2]})
- # Plot 1: Velocity
- ax1.plot(time_axis, velocity_cm_s, 'k-', linewidth=1.5, label='Velocity')
- ax1.axvline(x=150/framerate, color='red', linestyle='--', linewidth=2, label='Looming start')
- ax1.set_ylabel('Velocity (cm/s)', fontsize=12, fontweight='bold')
- ax1.set_title(f'Trial: {t_id} ({m_id}) - Group: {g}, Rank: {r}', fontsize=14, fontweight='bold')
- ax1.legend(loc='upper right')
- ax1.grid(True, alpha=0.3)
- ax1.set_ylim(bottom=0)
- # Plot 2: Ethogram
- y_pos = 0
- behavior_positions = {}
- for behavior in behaviors:
- behavior_positions[behavior] = y_pos
- color = behavior_colors.get(behavior, 'gray')
- # Get behavior frames (relative to plot window)
- for frame_idx in range(plot_start_frame, plot_end_frame):
- relative_frame = frame_idx - looming_start_frame + 150 # Relative to looming window
- try:
- if relative_frame in trial_bhvr_data[behavior].index:
- if trial_bhvr_data[behavior][relative_frame] == 1:
- time_point = (frame_idx - plot_start_frame) / framerate
- ax2.barh(y_pos, 1/framerate, left=time_point, height=0.8,
- color=color, edgecolor='none', alpha=0.8)
- except:
- pass
- y_pos += 1
- # Add vertical line for looming start
- ax2.axvline(x=150/framerate, color='red', linestyle='--', linewidth=2, label='Looming start')
- ax2.set_yticks(range(len(behaviors)))
- ax2.set_yticklabels(behaviors, fontsize=10)
- ax2.set_xlabel('Time (s)', fontsize=12, fontweight='bold')
- ax2.set_ylabel('Behaviors', fontsize=12, fontweight='bold')
- ax2.set_xlim(0, (plot_end_frame - plot_start_frame) / framerate)
- ax2.grid(True, axis='x', alpha=0.3)
- # Add legend for behaviors
- legend_patches = [mpatches.Patch(color=behavior_colors.get(b, 'gray'), label=b)
- for b in behaviors if b in behavior_colors]
- ax2.legend(handles=legend_patches, loc='upper right', ncol=3, fontsize=9)
- plt.tight_layout()
- plt.savefig(f'data/ethogram_velocity_plots/lst{g}_{t_id}_ethogram_velocity.png', dpi=300, bbox_inches='tight')
- plt.close()
- plot_count += 1
- if plot_count % 10 == 0:
- print(f"Processed {plot_count} trials...")
- except Exception as e:
- print(f"Error processing trial {t_id}: {str(e)}")
- continue
- print(f"\nCompleted! Generated {plot_count} plots.")
- print(f"All plots saved in: data/ethogram_velocity_plots/")
- print(f"Time window for each plot: -5s to +60s relative to looming start")
- # %% [markdown]
- # ### ret: plot ethogram
- # %%
- # Plot ethogram for ret trials with x-coordinate trajectory
- import matplotlib.pyplot as plt
- import matplotlib.patches as mpatches
- import os
- import pandas as pd
- import numpy as np
- from analyze_data_utils import ret_extract_target_data
- # Create output directory if it doesn't exist
- os.makedirs('data/ret_ethogram_plots', exist_ok=True)
- # Extract x-coordinate data for all trials
- x_coordinate_data = ret_extract_target_data(ret_raw_data, target_key=['X center'], times=['withrat'])
- # Behavior colors for ret
- behavior_colors_ret = {
- 'rat_in': 'black',
- 'approach': 'red',
- 'investigation': 'orange',
- 'withdrawal': 'blue',
- 'stretch-attend': 'cyan',
- 'freezing': 'navy',
- 'tail_rattling': 'purple',
- 'huddling': 'pink',
- 'approach_partner': 'lightcoral',
- 'follow_partner': 'salmon',
- 'groom_partner': 'gold',
- 'sniff_partner': 'yellow',
- 'grooming': 'green'
- }
- # Loop through all trials
- plot_count = 0
- for g in ['s', 'p']:
- for r in ['D', 'S']:
- if g not in ret_bhvr_frames_dict or r not in ret_bhvr_frames_dict[g]:
- continue
- for t_id in ret_bhvr_frames_dict[g][r].keys():
- # Get behavior data for this trial
- trial_bhvr_data = ret_bhvr_frames_dict[g][r][t_id]
- if g == 's':
- if t_id[1] in ['4']:
- if t_id[-1] == 'S':
- m_id = '1'
- elif t_id[-1] == 'D':
- m_id = '2'
- else:
- if t_id[-1] == 'D':
- m_id = '1'
- elif t_id[-1] == 'S':
- m_id = '2'
- sess_id = t_id.split('M')[0]+'M'+m_id
- elif g =='p':
- sess_id = t_id.split('M')[0]+'M1&M2'
- # Extract x-coordinate data for this session
- session_label = t_id.split('M')[0] # e.g., 'D1'
- cond = 'MD' if r == 'D' else 'MS'
- # Get x-coordinate series
- if g == 's':
- x_series = x_coordinate_data['single'][cond]['withrat'][session_label]
- else:
- x_series = x_coordinate_data['pair'][cond]['withrat'][session_label]
- # Convert to numeric and handle any non-numeric values
- x_values = pd.to_numeric(x_series, errors='coerce')
- # Calculate the offset between behavior and x-coordinate data
- offset_frames = rat_in_frames[sess_id] - offset_frames[sess_id]
- # Behavior data range (always 0 to 8900)
- behavior_start = 0
- behavior_end = 8900
- # Extract x-coordinate data with offset handling
- if offset_frames >= 0:
- # Normal case: x-coordinate starts before or at behavior start
- plot_x_values = x_values[offset_frames:offset_frames + 8900]
- x_time_offset = 0
- else:
- # offset_frames is negative: behavior starts before x-coordinate data
- # Need to pad the beginning with NaN and shift x data to the right
- x_start = 0
- x_end = 8900 + offset_frames # This will be less than 8900
- plot_x_values = pd.Series([np.nan] * (-offset_frames)).append(
- x_values[x_start:x_end], ignore_index=True)
- x_time_offset = -offset_frames
- # Get all behaviors
- behaviors = [b for b in trial_bhvr_data.keys() if b != 'rat_in']
- # Create figure with two subplots (x-coordinate and ethogram)
- fig, (ax1, ax2) = plt.subplots(2, 1, figsize=(16, 10), sharex=True,
- gridspec_kw={'height_ratios': [1, 2]})
- # Plot 1: X-coordinate trajectory
- time_axis = np.arange(len(plot_x_values)) / framerate
- ax1.plot(time_axis, plot_x_values, color='black', linewidth=0.5, alpha=0.7)
- ax1.set_ylabel('X Coordinate (cm)', fontsize=12, fontweight='bold')
- ax1.set_title(f'RET Trial: {t_id} - Group: {g}, Rank: {r}', fontsize=14, fontweight='bold')
- ax1.grid(True, axis='both', alpha=0.3)
- # Add vertical lines for time markers on x-coordinate plot
- ax1.axvline(x=0/framerate, color='red', linestyle='--', linewidth=2, alpha=0.5)
- ax1.axvline(x=1800/framerate, color='red', linestyle='--', linewidth=2, alpha=0.5)
- ax1.axvline(x=3600/framerate, color='red', linestyle='--', linewidth=2, alpha=0.5)
- ax1.axvline(x=5400/framerate, color='red', linestyle='--', linewidth=2, alpha=0.5)
- ax1.axvline(x=7200/framerate, color='red', linestyle='--', linewidth=2, alpha=0.5)
- ax1.axvline(x=8900/framerate, color='red', linestyle='--', linewidth=2, alpha=0.5)
- # Plot 2: Ethogram
- y_pos = 0
- behavior_positions = {}
- for behavior in behaviors:
- behavior_positions[behavior] = y_pos
- color = behavior_colors_ret.get(behavior, 'gray')
- # Get behavior frames (behavior data is always from 0 to 8900)
- for frame_idx in range(behavior_start, behavior_end):
- try:
- if frame_idx in trial_bhvr_data[behavior].index:
- if trial_bhvr_data[behavior][frame_idx] == 1:
- time_point = frame_idx / framerate
- ax2.barh(y_pos, 1/framerate, left=time_point, height=0.8,
- color=color, edgecolor='none', alpha=0.8)
- except:
- pass
- y_pos += 1
- # Add vertical lines for time markers on ethogram
- ax2.axvline(x=0/framerate, color='red', linestyle='--', linewidth=2)
- ax2.axvline(x=1800/framerate, color='red', linestyle='--', linewidth=2)
- ax2.axvline(x=3600/framerate, color='red', linestyle='--', linewidth=2)
- ax2.axvline(x=5400/framerate, color='red', linestyle='--', linewidth=2)
- ax2.axvline(x=7200/framerate, color='red', linestyle='--', linewidth=2)
- ax2.axvline(x=8900/framerate, color='red', linestyle='--', linewidth=2)
- ax2.set_yticks(range(len(behaviors)))
- ax2.set_yticklabels(behaviors, fontsize=10)
- ax2.set_xlabel('Time (s)', fontsize=12, fontweight='bold')
- ax2.set_ylabel('Behaviors', fontsize=12, fontweight='bold')
- ax2.set_xlim(0, 8900 / framerate)
- ax2.grid(True, axis='x', alpha=0.3)
- # Add legend for behaviors
- legend_patches = [mpatches.Patch(color=behavior_colors_ret.get(b, 'gray'), label=b)
- for b in behaviors if b in behavior_colors_ret]
- # ax2.legend(handles=legend_patches, loc='upper right', ncol=4, fontsize=8)
- plt.tight_layout()
- plt.savefig(f'data/ret_ethogram_plots/ret{g}_{t_id}_ethogram_with_x.png', dpi=300, bbox_inches='tight')
- plt.close()
- plot_count += 1
- if plot_count % 10 == 0:
- print(f"Processed {plot_count} trials...")
- print(f"\nCompleted! Generated {plot_count} plots.")
- print(f"All plots saved in: data/ret_ethogram_plots/")
- print(f"Time window for each plot: 0 - 5 min")
- # %% [markdown]
- # ## Fig1
- # %% [markdown]
- # ### Fig1B&D. location and velocity histogram
- # %%
- import numpy as np
- import matplotlib.pyplot as plt
- from matplotlib.colors import ListedColormap, BoundaryNorm
- import seaborn as sns
- from analyze_data_utils import filter_in_range, filter_dict_data
- import matplotlib.colors as mcolors
- import pandas as pd
- nbins = 25
- cmap = 'viridis'
- pixpergrid = 43.52
- pixpercm = pixpergrid / 2 # pixels per cm
- def prepare_velocity_data(ids, data_source, min_speed=0, max_speed=300):
- all_x = []
- all_y = []
- all_vx = []
- all_vy = []
- all_ids = []
- for mouse_id in ids:
- # interested_frame = (looming_start_frames[mouse_id][0], looming_start_frames[mouse_id][-1]+1800)
- # interested_frame = [0, 18000]
- # x = data_source[mouse_id]['coordinate']['waist']['x'][interested_frame[0]:interested_frame[1]]
- # y = data_source[mouse_id]['coordinate']['waist']['y'][interested_frame[0]:interested_frame[1]]
- # # Calculate velocity vectors (cm/s)
- # vx = np.diff(x) / pixpercm * framerate
- # vy = np.diff(y) / pixpercm * framerate
- # # Keep only points with speed >= min_speed
- # speed = np.sqrt(vx**2 + vy**2)
- # mask = (speed >= min_speed) & (speed <= max_speed)
- # all_x.extend(x[:-1][mask]) # x[:-1] aligns with diff length
- # all_y.extend(y[:-1][mask])
- # all_vx.extend(vx[mask])
- # all_vy.extend(vy[mask])
- # all_ids.extend([mouse_id] * np.sum(mask)) # Record mouse ID for each point
- start_frames = looming_start_frames[mouse_id]
- for start_frame in start_frames:
- start_frame = start_frame + 30 * 20
- end_frame = start_frame + 30 * 60
- max_frame = len(data_source[mouse_id]['coordinate']['waist']['x'])
- end_frame = min(end_frame, max_frame)
- x = data_source[mouse_id]['coordinate']['waist']['x'][start_frame:end_frame]
- y = data_source[mouse_id]['coordinate']['waist']['y'][start_frame:end_frame]
- vx = np.diff(x) / pixpercm * framerate
- vy = np.diff(y) / pixpercm * framerate
- speed = np.sqrt(vx**2 + vy**2)
- mask = (speed >= min_speed) & (speed <= max_speed)
- all_x.extend(x[:-1][mask])
- all_y.extend(y[:-1][mask])
- all_vx.extend(vx[mask])
- all_vy.extend(vy[mask])
- all_ids.extend([mouse_id] * np.sum(mask))
- return np.array(all_x), np.array(all_y), np.array(all_vx), np.array(all_vy), np.array(all_ids)
- # def prepare_velocity_data(ids, data_source, min_speed=0, max_speed=300):
- # all_x = []
- # all_y = []
- # all_vx = []
- # all_vy = []
- # all_ids = []
- # for mouse_id in ids:
- # # Get the entire time period [first stimulus start, last stimulus end+1800]
- # start_frame0 = looming_start_frames[mouse_id][0]
- # end_frame_last = looming_start_frames[mouse_id][-1] + 1800
- # max_frame = len(data_source[mouse_id]['coordinate']['waist']['x'])
- # end_frame_last = min(end_frame_last, max_frame) # Ensure not exceeding data range
- # # Extract coordinate data for the entire time period
- # x_full = data_source[mouse_id]['coordinate']['waist']['x'][start_frame0:end_frame_last]
- # y_full = data_source[mouse_id]['coordinate']['waist']['y'][start_frame0:end_frame_last]
- # # Create a mask initialized to True, indicating all points are initially included
- # mask = np.ones(len(x_full), dtype=bool)
- # # Iterate through each stimulus event, excluding data during events
- # for start_frame in looming_start_frames[mouse_id]:
- # # Calculate the relative position of the current event in the full data
- # event_start = start_frame - start_frame0
- # event_end = min(event_start + 1800, len(x_full))
- # # Mark data during the event as False (excluded)
- # if event_start < len(mask):
- # mask[event_start:event_end] = False
- # # Apply mask to exclude all data during stimulus events
- # x = x_full[mask]
- # y = y_full[mask]
- # # Calculate velocity vectors (cm/s)
- # vx = np.diff(x) / pixpercm * framerate
- # vy = np.diff(y) / pixpercm * framerate
- # # Keep only points with speed in [min_speed, max_speed]
- # speed = np.sqrt(vx**2 + vy**2)
- # speed_mask = (speed >= min_speed) & (speed <= max_speed)
- # # Add data to total lists
- # all_x.extend(x[:-1][speed_mask]) # x[:-1] aligns with diff length
- # all_y.extend(y[:-1][speed_mask])
- # all_vx.extend(vx[speed_mask])
- # all_vy.extend(vy[speed_mask])
- # all_ids.extend([mouse_id] * np.sum(speed_mask))
- # return np.array(all_x), np.array(all_y), np.array(all_vx), np.array(all_vy), np.array(all_ids)
- def prepare_nest_velocity(ids, data_source, nest_data, min_speed=5, max_speed=300):
- all_x = []
- all_y = []
- all_vx = []
- all_vy = []
- all_ids = []
- for mouse_id in ids:
- interested_frame = (looming_start_frames[mouse_id][0], looming_start_frames[mouse_id][-1]+1800)
- # interested_frame = [0, 18000]
- nest_id = mouse_id.split('M')[0]+'M'+mouse_id[-2]+'T1'
- g = 's'
- r = mouse_id[-2]
- start_frames = looming_start_frames[mouse_id]
- for start_frame in start_frames:
- # start_frame = 9000
- if nest_id in nest_data[g][r].keys():
- nest_list = nest_data[g][r][nest_id]
- else:
- continue
- for nest in nest_list:
- if np.any(np.isnan(nest)):
- continue
- else:
- start, end = nest
- nest_start_frame = start_frame + start
- # interested_frame = [nest_start_frame-60, nest_start_frame+10] # baseline nest return
- interested_frame = [nest_start_frame-150-60, nest_start_frame-150+10] # looming nest return
- # interested_frame = [start_frame+start-150-10, start_frame+end-150+10] # escape
- x = data_source[mouse_id]['coordinate']['waist']['x'][interested_frame[0]:interested_frame[1]]
- y = data_source[mouse_id]['coordinate']['waist']['y'][interested_frame[0]:interested_frame[1]]
- # Calculate velocity vectors (cm/s)
- vx = np.diff(x) / pixpercm * framerate
- vy = np.diff(y) / pixpercm * framerate
- # Keep only points with speed >= min_speed
- speed = np.sqrt(vx**2 + vy**2)
- mask = (speed >= min_speed) & (speed <= max_speed)
- all_x.extend(x[:-1][mask]) # x[:-1] aligns with diff length
- all_y.extend(y[:-1][mask])
- all_vx.extend(vx[mask])
- all_vy.extend(vy[mask])
- all_ids.extend([mouse_id] * np.sum(mask)) # Record mouse ID for each point
- return np.array(all_x), np.array(all_y), np.array(all_vx), np.array(all_vy), np.array(all_ids)
- # all_speed_avg = []
- # for m_id in ms_ids:
- # nest_data = filter_dict_data(baseline_behavior_frames_dict, 'nest')
- nest_data = filter_dict_data(lst_baseline_bhvr_labels_dict, 'nest')
- nest_data = filter_in_range(nest_data, t_range=(0,18000), method='replace')
- # x, y, vx, vy, point_ids = prepare_nest_velocity(ms_ids, lsts_data, nest_data)
- x, y, vx, vy, point_ids = prepare_velocity_data(ms_ids, lsts_data)
- heatmap, xedges, yedges = np.histogram2d(x, y, range=[[0, 1088], [0, 1088]], bins=nbins)
- mouse_count_grid = np.zeros_like(heatmap, dtype=int)
- x_bin_indices = np.digitize(x, xedges) - 1
- y_bin_indices = np.digitize(y, yedges) - 1
- cell_mice_dict = {}
- for xi, yi, mouse_id in zip(x_bin_indices, y_bin_indices, point_ids):
- if 0 <= xi < nbins and 0 <= yi < nbins:
- key = (xi, yi)
- if key not in cell_mice_dict:
- cell_mice_dict[key] = set()
- cell_mice_dict[key].add(mouse_id)
- for (xi, yi), mice_set in cell_mice_dict.items():
- mouse_count_grid[xi, yi] = len(mice_set)
- # mask = mouse_count_grid >= 3
- grid_vx, _, _ = np.histogram2d(x, y, range=[[0, 1088], [0, 1088]], bins=nbins, weights=vx)
- grid_vy, _, _ = np.histogram2d(x, y, range=[[0, 1088], [0, 1088]], bins=nbins, weights=vy)
- count, _, _ = np.histogram2d(x, y, range=[[0, 1088], [0, 1088]], bins=nbins)
- # count[count == 0] = 1
- grid_vx /= count
- grid_vy /= count
- grid_vx /= 2
- grid_vy /= 2 # change scale for nest return and escape
- mask = (mouse_count_grid >= 3) & (count >= 5)
- grid_vx[~mask] = 1e-5
- grid_vy[~mask] = 1e-5
- speed_grid = np.sqrt(grid_vx**2 + grid_vy**2)
- speeds = np.sqrt(vx**2 + vy**2)
- speed_sum, _, _ = np.histogram2d(x, y, range=[[0, 1088], [0, 1088]], bins=[xedges, yedges], weights=speeds)
- speed_avg = speed_sum / count
- speed_avg[~mask] = 1e-5
- # all_speed_avg.append(speed_avg.T)
- X, Y = np.meshgrid((xedges[:-1] + xedges[1:]) / 2,
- (yedges[:-1] + yedges[1:]) / 2)
- hist_frequency = heatmap / np.sum(heatmap) * 100
- hist_frequency[~mask] = 1e-5
- log_heatmap = np.log1p(hist_frequency)
- magnitude = np.sqrt(grid_vx**2 + grid_vy**2)
- scaled_magnitude = np.sqrt(magnitude + 1e-5)
- vx_mod = grid_vx / (magnitude + 1e-5) * scaled_magnitude
- vy_mod = grid_vy / (magnitude + 1e-5) * scaled_magnitude
- # Save data for plotting
- df = pd.DataFrame({
- 'x_index': np.repeat(np.arange(nbins), nbins),
- 'y_index': np.tile(np.arange(nbins), nbins),
- 'time_pct': hist_frequency.T.flatten(),
- 'speed_avg': speed_avg.T.flatten(),
- 'vx': grid_vx.T.flatten(),
- 'vy': grid_vy.T.flatten()
- })
- df.to_csv('data/Fig1B&D_location_speed_heatmap_60s.csv', index=False)
- # %%
- import numpy as np
- import matplotlib.pyplot as plt
- from matplotlib.colors import ListedColormap, BoundaryNorm
- import seaborn as sns
- import matplotlib.colors as mcolors
- import pandas as pd
- nbins = 25
- df = pd.read_csv('data/Fig1B&D_location_speed_heatmap_60s.csv')
- # Reshape to an (nbins, nbins) matrix (no transpose, matching the shape of the original computed variable).
- hist_frequency = df.pivot(index='y_index', columns='x_index', values='time_pct').values
- speed_avg = df.pivot(index='y_index', columns='x_index', values='speed_avg').values
- grid_vx = df.pivot(index='y_index', columns='x_index', values='vx').values
- grid_vy = df.pivot(index='y_index', columns='x_index', values='vy').values
- xedges = np.linspace(0, nbins, nbins+1)
- yedges = xedges.copy()
- fig, ax = plt.subplots(figsize=(5,5), dpi=200)
- # 1) draw coordinate heatmap
- im = ax.imshow(
- hist_frequency.T,
- cmap=cmap,
- origin='lower', # Start drawing from bottom-left corner
- extent=(0, nbins, 0, nbins), # Map x and y axes to [0, nbins]
- vmin=0,
- vmax=4
- )
- fig.colorbar(im, label='Time (%)')
- line_cor = 0.3
- 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)
- ax.plot([0+line_cor, 0+line_cor], [17.5*nbins/25-line_cor, 25*nbins/25-line_cor], color='white', linewidth=2)
- ax.plot([0+line_cor, 10*nbins/25+line_cor], [25*nbins/25-line_cor, 25*nbins/25-line_cor], color='white', linewidth=2)
- ax.set_xlim(0, nbins)
- ax.set_ylim(0, nbins)
- ax.set_xticks([])
- ax.set_yticks([])
- ax.invert_yaxis()
- plt.tight_layout()
- # plt.savefig(f'G:/Li_lab/ppt/S_paper/Paper_v7/new_figures/coord_{title}.eps', format="eps", dpi=300, bbox_inches="tight")
- plt.show()
- fig, ax = plt.subplots(figsize=(5,5), dpi=300)
- # 1) draw speed heatmap
- plt.imshow(speed_avg.T, cmap=cmap, origin='lower',
- extent=(0, nbins, 0, nbins), vmin=0, vmax=40)
- plt.colorbar(label='Speed (cm/s)')
- # 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)
- # ax.plot([0+line_cor, 0+line_cor], [17.5*nbins/25-line_cor, 25*nbins/25-line_cor], color='white', linewidth=2)
- # ax.plot([0+line_cor, 10*nbins/25+line_cor], [25*nbins/25-line_cor, 25*nbins/25-line_cor], color='white', linewidth=2)
- # ax.set_xticks([])
- # ax.set_yticks([])
- # ax.invert_yaxis()
- # 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")
- # plt.show()
- # fig, ax = plt.subplots(figsize=(5,5), dpi=200)
- # ax.set_facecolor('black') # Plot area background
- X, Y = np.meshgrid(np.arange(nbins) + 0.5, np.arange(nbins) + 0.5)
- scale = 0.5
- head_scale = 0.1
- ax.set_xlim(X.min(), X.max())
- ax.set_ylim(Y.min(), Y.max())
- for i in range(0, X.shape[0], 1):
- for j in range(0, X.shape[1], 1):
- vx_ = grid_vx.T[i, j]
- vy_ = grid_vy.T[i, j]
- mag = np.sqrt(vx_**2 + vy_**2)
- dx = vx_ * scale
- dy = vy_ * scale
- # x0 = X[i, j] - dx / 2
- # y0 = Y[i, j] - dy / 2
- # x1 = X[i, j] + dx / 2
- # y1 = Y[i, j] + dy / 2
- x0 = X[i, j]
- y0 = Y[i, j]
- x1 = x0 + dx
- y1 = y0 + dy
- ax.annotate('', xy=(x1, y1), xytext=(x0, y0),
- arrowprops=dict(
- arrowstyle='->, head_width={:.2f}, head_length={:.2f}'.format(mag * head_scale, mag * head_scale * 1.5),
- color='white',
- linewidth=1,
- mutation_scale=5,
- shrinkA=0, shrinkB=0
- ))
- ax.invert_yaxis()
- line_cor = 0.3
- 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)
- ax.plot([0+line_cor, 0+line_cor], [17.5*nbins/25-line_cor, 25*nbins/25-line_cor], color='white', linewidth=2)
- ax.plot([0+line_cor, 10*nbins/25+line_cor], [25*nbins/25-line_cor, 25*nbins/25-line_cor], color='white', linewidth=2)
- ax.set_xlim(0, nbins)
- ax.set_ylim(0, nbins)
- ax.set_xticks([])
- ax.set_yticks([])
- ax.invert_yaxis()
- plt.tight_layout()
- # plt.savefig(f'G:/Li_lab/ppt/S_paper/Paper_v7/new_figures/speed_{title}.eps', format="eps", dpi=300, bbox_inches="tight")
- plt.show()
- # %% [markdown]
- # ### Fig1C&E & S1B: Time in zone, Speed in zone.
- # %%
- import numpy as np
- import pandas as pd
- from analyze_data_utils import get_lst_location_value
- # Define parameters
- framerate = 30 # fps
- pixpergrid = 43.52
- pixpercm = pixpergrid / 2 # pixels per cm
- # Define time periods (in frames)
- time_periods = {
- 'baseline': (0, 10*60*framerate), # baseline 10 minutes
- '0-5s': (0, 5 * framerate), # 0-5 seconds after looming
- '5-20s': (5 * framerate, 20 * framerate), # 5-20 seconds after looming
- '20-60s': (20 * framerate, 60 * framerate) # 20-60 seconds after looming
- }
- # Define zone labels
- zone_labels = {0: 'nest', 1: 'edge', 2: 'center'}
- # zones_time_pcts[group][rank][time_period][zone] = [percentages for each mouse]
- zones_time_pcts = {'s': {'D': {}, 'S': {}}, 'p': {'D': {}, 'S': {}}}
- zones_speed = {'s': {'D': {}, 'S': {}}, 'p': {'D': {}, 'S': {}}}
- for g in ['s', 'p']:
- for m in ['D', 'S']:
- for period in time_periods.keys():
- zones_time_pcts[g][m][period] = {0: [], 1: [], 2: []} # nest, edge, center
- zones_speed[g][m][period] = {0: [], 1: [], 2: []}
- for g in ['s', 'p']:
- group_name = 'single' if g == 's' else 'pair'
- for m in ['D', 'S']:
- rank_name = 'Dominant' if m == 'D' else 'Subordinate'
- m_ids = mS_ids if m == 'S' else mD_ids
- lst_data = lsts_data if g == 's' else lstp_data
- for m_id in m_ids:
- for period_name, (start_offset, end_offset) in time_periods.items():
- if period_name == 'baseline':
- interested_frame = (start_offset, end_offset)
- coord_data = {
- 'x': lst_data[m_id]['coordinate']['waist']['x'][interested_frame[0]:interested_frame[1]].reset_index(drop=True),
- 'y': lst_data[m_id]['coordinate']['waist']['y'][interested_frame[0]:interested_frame[1]].reset_index(drop=True)
- }
- location_value = get_lst_location_value(coord_data)
- velocity_value = lst_data[m_id]['velocity']['waist'][interested_frame[0]:interested_frame[1]].reset_index(drop=True)
- for zone in [0, 1, 2]:
- zone_count = (location_value == zone).sum()
- zone_pct = zone_count / len(location_value) * 100
- zones_time_pcts[g][m][period_name][zone].append(zone_pct)
- zone_speed = velocity_value[location_value == zone].mean() / pixpercm * framerate
- zones_speed[g][m][period_name][zone].append(zone_speed)
- else:
- for lsf in looming_start_frames.get(m_id, []):
- interested_frame = (lsf + start_offset, lsf + end_offset)
- coord_data = {
- 'x': lst_data[m_id]['coordinate']['waist']['x'][interested_frame[0]:interested_frame[1]].reset_index(drop=True),
- 'y': lst_data[m_id]['coordinate']['waist']['y'][interested_frame[0]:interested_frame[1]].reset_index(drop=True)
- }
- location_value = get_lst_location_value(coord_data)
- velocity_value = lst_data[m_id]['velocity']['waist'][interested_frame[0]:interested_frame[1]].reset_index(drop=True)
- for zone in [0, 1, 2]:
- zone_count = (location_value == zone).sum()
- zone_pct = zone_count / len(location_value) * 100
- zones_time_pcts[g][m][period_name][zone].append(zone_pct)
- zone_speed = velocity_value[location_value == zone].mean() / pixpercm * framerate
- zones_speed[g][m][period_name][zone].append(zone_speed)
- # %%
- # Convert zones_time_pcts to DataFrame (long format)
- # Each row represents one mouse, first all Dom, then all Sub
- import pandas as pd
- import numpy as np
- # Define order of time periods and zones
- time_period_order = ['baseline', '0-5s', '5-20s', '20-60s']
- zone_order = ['nest', 'edge', 'center']
- df_zones_pct = {'s': None, 'p': None}
- for g in ['s', 'p']:
- data_list = []
- for m in ['D', 'S']: # D (Dom) first, then S (Sub)
- rank_name = 'Dom' if m == 'D' else 'Sub'
- m_ids = mD_ids if m == 'D' else mS_ids
- for m_id in m_ids:
- # Extract B ID (e.g., 'B1M1(S)' -> 'B1')
- b_id = m_id.split('M')[0]
- # Initialize data row for this mouse
- row_data = {
- 'B_id': b_id,
- 'rank': rank_name,
- 'mouse_id': m_id,
- }
- # Iterate through each time period and zone
- for period in time_period_order:
- for zone_id, zone_name in enumerate(zone_order):
- # Column name format: zone_period (e.g., nest_baseline)
- col_name = f"{zone_name}_{period}"
- # Get all data for this mouse in this time period and zone
- values = zones_time_pcts[g][m][period][zone_id]
- if period == 'baseline':
- # Baseline: each mouse has only one value
- # Find the index of this mouse in m_ids
- mouse_idx = m_ids.index(m_id)
- if mouse_idx < len(values):
- row_data[col_name] = values[mouse_idx]
- else:
- row_data[col_name] = np.nan
- else:
- # Other time periods: calculate average across all looming events for this mouse
- m_ids_list = mD_ids if m == 'D' else mS_ids
- mouse_idx = m_ids_list.index(m_id)
- # Calculate the number of looming events for this mouse
- looming_count = len(looming_start_frames.get(m_id, []))
- # Starting index for this mouse's data
- start_idx = sum([len(looming_start_frames.get(mid, []))
- for mid in m_ids_list[:mouse_idx]])
- end_idx = start_idx + looming_count
- # Get all looming data for this mouse
- mouse_values = values[start_idx:end_idx]
- # Calculate average
- if len(mouse_values) > 0:
- row_data[col_name] = np.nanmean(mouse_values)
- else:
- row_data[col_name] = np.nan
- # Add to list
- data_list.append(row_data)
- # Convert to DataFrame
- df_zones_pct[g] = pd.DataFrame(data_list)
- # Adjust column order
- # Information columns
- info_cols = ['B_id', 'rank', 'mouse_id']
- # Data columns
- data_cols = [f"{zone}_{period}" for period in time_period_order for zone in zone_order]
- for g in ['s', 'p']:
- df_zones_pct[g] = df_zones_pct[g][info_cols + data_cols]
- # Save DataFrame
- df_zones_pct['s'].to_csv('data/FigS1B_zones_time_pcts.csv', index=False)
- # %%
- df = pd.read_csv('data/FigS1B_zones_time_pcts.csv')
- print(df.columns.tolist())
- exclude_cols = ['B_id', 'rank', 'mouse_id']
- plot_cols = [col for col in df.columns if col not in exclude_cols]
- means = df[plot_cols].mean()
- plt.figure()
- bar_colors = ['green', 'blue', 'orange']
- bars = plt.bar(means.index, means.values, color=bar_colors[:len(means)])
- df_melted = df[plot_cols]
- sns.stripplot(
- data=df_melted,
- jitter=True,
- color='gray',
- size=5
- )
- plt.ylabel('Time in zone (%)')
- plt.xticks(rotation=45, ha='right')
- plt.show()
- # %%
- # Normalize data: convert from percentage to unit (s/(min∙cm²))
- import pandas as pd
- import numpy as np
- # Define parameters
- pixpergrid = 43.52
- pixpercm = pixpergrid / 2 # pixels per cm
- arena_size = 1088 # pixels
- # Define time period durations (minutes)
- time_durations_min = {
- 'baseline': 150 / 60, # 2.5 min
- '0-5s': 5 / 60, # 0.0833 min
- '5-20s': 15 / 60, # 0.25 min
- '20-60s': 40 / 60 # 0.667 min
- }
- # Define zone areas (cm²)
- zone_areas_cm2 = {
- 'nest': 80*4,
- 'edge': 280*4,
- 'center': 265*4
- }
- # Copy df_zones_wide
- df_zones_unit = df_zones_pct['s'].copy()
- # Normalize each data column
- for col in data_cols:
- # Parse column name: zone_period
- parts = col.split('_')
- if len(parts) >= 2:
- zone = parts[0] # nest, edge, center
- period = '_'.join(parts[1:]) # baseline, 0-5s, 5-20s, 20-60s
- # Get time duration and zone area
- time_min = time_durations_min.get(period, 1)
- area_cm2 = zone_areas_cm2.get(zone, 1)
- # Normalize: percentage / 100 / time(min) / area(cm²)
- # Multiply by 60 because the result unit is s/(min∙cm²),
- # converting percentage to decimal represents the proportion,
- # proportion * total time(s) = actual time(s)
- # actual time(s) / time(min) / area(cm²) = s/(min∙cm²)
- df_zones_unit[col] = df_zones_unit[col] / 100 * 60 / area_cm2
- # Save normalized data
- df_zones_unit.to_csv('data/Fig1C_zones_time_unit.csv', index=False)
- # %%
- df = pd.read_csv('data/Fig1C_zones_time_unit.csv')
- print(df.columns.tolist())
- exclude_cols = ['B_id', 'rank', 'mouse_id']
- plot_cols = [col for col in df.columns if col not in exclude_cols]
- means = df[plot_cols].mean()
- plt.figure()
- bar_colors = ['green', 'blue', 'orange']
- bars = plt.bar(means.index, means.values, color=bar_colors[:len(means)])
- df_melted = df[plot_cols]
- sns.stripplot(
- data=df_melted,
- jitter=True,
- color='gray',
- size=5
- )
- plt.ylabel('Time in zone (s/(min﹡cm²))')
- plt.xticks(rotation=45, ha='right')
- plt.show()
- # %%
- # Convert zones_speed to DataFrame (long format)
- # Each row represents one mouse, first all Dom, then all Sub
- import pandas as pd
- import numpy as np
- # Define order of time periods and zones
- time_period_order = ['baseline', '0-5s', '5-20s', '20-60s']
- zone_order = ['nest', 'edge', 'center']
- df_zones_speed = {'s': None, 'p': None}
- for g in ['s', 'p']:
- data_list = []
- for m in ['D', 'S']: # D (Dom) first, then S (Sub)
- rank_name = 'Dom' if m == 'D' else 'Sub'
- m_ids = mD_ids if m == 'D' else mS_ids
- for m_id in m_ids:
- # Extract B ID (e.g., 'B1M1(S)' -> 'B1')
- b_id = m_id.split('M')[0]
- # Initialize data row for this mouse
- row_data = {
- 'B_id': b_id,
- 'rank': rank_name,
- 'mouse_id': m_id,
- }
- # Iterate through each time period and zone
- for period in time_period_order:
- for zone_id, zone_name in enumerate(zone_order):
- # Column name format: zone_period (e.g., nest_baseline)
- col_name = f"{zone_name}_{period}"
- # Get all data for this mouse in this time period and zone
- values = zones_speed[g][m][period][zone_id]
- if period == 'baseline':
- # Baseline: each mouse has only one value
- # Find the index of this mouse in m_ids
- mouse_idx = m_ids.index(m_id)
- if mouse_idx < len(values):
- row_data[col_name] = values[mouse_idx]
- else:
- row_data[col_name] = np.nan
- else:
- # Other time periods: calculate average across all looming events for this mouse
- m_ids_list = mD_ids if m == 'D' else mS_ids
- mouse_idx = m_ids_list.index(m_id)
- # Calculate the number of looming events for this mouse
- looming_count = len(looming_start_frames.get(m_id, []))
- # Starting index for this mouse's data
- start_idx = sum([len(looming_start_frames.get(mid, []))
- for mid in m_ids_list[:mouse_idx]])
- end_idx = start_idx + looming_count
- # Get all looming data for this mouse
- mouse_values = values[start_idx:end_idx]
- # Calculate average
- if len(mouse_values) > 0:
- row_data[col_name] = np.nanmean(mouse_values)
- else:
- row_data[col_name] = np.nan
- # Add to list
- data_list.append(row_data)
- # Convert to DataFrame
- df_zones_speed[g] = pd.DataFrame(data_list)
- # Adjust column order
- # Information columns
- info_cols = ['B_id', 'rank', 'mouse_id']
- # Data columns
- data_cols = [f"{zone}_{period}" for period in time_period_order for zone in zone_order]
- for g in ['s', 'p']:
- df_zones_speed[g] = df_zones_speed[g][info_cols + data_cols]
- # Save DataFrame
- df_zones_speed['s'].to_csv('data/Fig1E_zones_speed.csv', index=False)
- # %%
- df = pd.read_csv('data/Fig1E_zones_speed.csv')
- print(df.columns.tolist())
- exclude_cols = ['B_id', 'rank', 'mouse_id']
- plot_cols = [col for col in df.columns if col not in exclude_cols]
- means = df[plot_cols].mean()
- plt.figure()
- bar_colors = ['green', 'blue', 'orange']
- bars = plt.bar(means.index, means.values, color=bar_colors[:len(means)])
- df_melted = df[plot_cols]
- sns.stripplot(
- data=df_melted,
- jitter=True,
- color='gray',
- size=5
- )
- plt.ylabel('Speed in zone (cm/s)')
- plt.xticks(rotation=45, ha='right')
- plt.show()
- # %% [markdown]
- # ### Fig1F. ΔTime in zone (%)
- # %%
- import pandas as pd
- import numpy as np
- # Define order of time periods and zones
- time_period_order = ['baseline', '0-5s', '5-20s', '20-60s']
- zone_order = ['nest', 'edge', 'center']
- for period in time_period_order:
- period_cols = [col for col in df_zones_pct['s'].columns if period in col]
- # Extract single and pair baseline data
- df_single = df_zones_pct['s'][['B_id', 'rank'] + period_cols].copy()
- df_pair = df_zones_pct['p'][['B_id', 'rank'] + period_cols].copy()
- # Merge data (by B_id and rank)
- df_merged = df_single.merge(df_pair, on=['B_id', 'rank'], suffixes=('_s', '_p'))
- # Calculate difference (pair - single)
- delta_data = []
- for b_id in df_merged['B_id'].unique():
- row = {'B_id': b_id}
- b_data = df_merged[df_merged['B_id'] == b_id]
- for zone in zone_order:
- # Dom data
- dom = b_data[b_data['rank'] == 'Dom']
- if len(dom) > 0:
- col_s = f'{zone}_{period}_s'
- col_p = f'{zone}_{period}_p'
- row[f'{zone}_D'] = dom[col_p].values[0] - dom[col_s].values[0]
- else:
- row[f'{zone}_D'] = np.nan
- # Sub data
- sub = b_data[b_data['rank'] == 'Sub']
- if len(sub) > 0:
- col_s = f'{zone}_{period}_s'
- col_p = f'{zone}_{period}_p'
- row[f'{zone}_S'] = sub[col_p].values[0] - sub[col_s].values[0]
- else:
- row[f'{zone}_S'] = np.nan
- delta_data.append(row)
- # Create DataFrame
- df_delta_time = pd.DataFrame(delta_data)
- # Adjust column order
- df_delta_time = df_delta_time[
- ['B_id', 'nest_D', 'nest_S', 'edge_D', 'edge_S', 'center_D', 'center_S']
- ]
- # Save results
- df_delta_time.to_csv(f'data/Fig1F_delta_zones_time_{period}.csv', index=False)
- # %%
- for period in time_period_order:
- df = pd.read_csv(f'data/Fig1F_delta_zones_time_{period}.csv')
- groups = ['nest', 'edge', 'center']
- fig, ax = plt.subplots()
- x_pos = {}
- xticks = []
- xticklabels = []
- x = 0
- colors = {'D': 'orange', 'S': 'blue'}
- # 1. x
- for g in groups:
- for cond in ['D', 'S']:
- if cond == 'D':
- x_pos[f'{g}_{cond}_mean'] = x
- x_pos[f'{g}_{cond}_pts'] = x + 0.8
- xticks += [x, x + 0.8]
- xticklabels += [f'{g}_{cond}_mean', f'{g}_{cond}_pts']
- else:
- x_pos[f'{g}_{cond}_pts'] = x
- x_pos[f'{g}_{cond}_mean'] = x + 0.8
- xticks += [x, x + 0.8]
- xticklabels += [f'{g}_{cond}_pts', f'{g}_{cond}_mean']
- x += 2.2
- # 2. mean ± SD
- for g in groups:
- for cond in ['D', 'S']:
- col = f'{g}_{cond}'
- xpos = x_pos[f'{g}_{cond}_mean']
- ax.errorbar(
- xpos,
- df[col].mean(),
- yerr=df[col].std(),
- fmt='o',
- color=colors[cond],
- capsize=4,
- markersize=9,
- elinewidth=1.5,
- capthick=1.5,
- zorder=3
- )
- # 3. raw points + pairing
- for i in range(len(df)):
- for g in groups:
- D_col = f'{g}_D'
- S_col = f'{g}_S'
- x_d = x_pos[f'{g}_D_pts']
- y_d = df.loc[i, D_col]
- x_s = x_pos[f'{g}_S_pts']
- y_s = df.loc[i, S_col]
- ax.scatter(x_d, y_d, color=colors['D'])
- ax.scatter(x_s, y_s, color=colors['S'])
- ax.plot([x_d, x_s], [y_d, y_s],
- color='gray', linewidth=1)
- # 4. axis
- ax.set_xticks(xticks)
- ax.set_xticklabels(xticklabels, rotation=45, ha='right')
- ax.set_ylabel('Δ Time in zone (%)')
- ax.set_ylim([-100, 100])
- ax.set_title(f'Δ Time in zone ({period})')
- plt.tight_layout()
- plt.show()
- # %% [markdown]
- # ### Fig1G. Δ Speed in zone (cm/s)
- # %%
- import pandas as pd
- import numpy as np
- # Define order of time periods and zones
- time_period_order = ['baseline', '0-5s', '5-20s', '20-60s']
- zone_order = ['nest', 'edge', 'center']
- for period in time_period_order:
- period_cols = [col for col in df_zones_speed['s'].columns if period in col]
- # Extract single and pair baseline data
- df_single = df_zones_speed['s'][['B_id', 'rank'] + period_cols].copy()
- df_pair = df_zones_speed['p'][['B_id', 'rank'] + period_cols].copy()
- # Merge data (by B_id and rank)
- df_merged = df_single.merge(df_pair, on=['B_id', 'rank'], suffixes=('_s', '_p'))
- # Calculate difference (pair - single)
- delta_data = []
- for b_id in df_merged['B_id'].unique():
- row = {'B_id': b_id}
- b_data = df_merged[df_merged['B_id'] == b_id]
- for zone in zone_order:
- # Dom data
- dom = b_data[b_data['rank'] == 'Dom']
- if len(dom) > 0:
- col_s = f'{zone}_{period}_s'
- col_p = f'{zone}_{period}_p'
- row[f'{zone}_D'] = dom[col_p].values[0] - dom[col_s].values[0]
- else:
- row[f'{zone}_D'] = np.nan
- # Sub data
- sub = b_data[b_data['rank'] == 'Sub']
- if len(sub) > 0:
- col_s = f'{zone}_{period}_s'
- col_p = f'{zone}_{period}_p'
- row[f'{zone}_S'] = sub[col_p].values[0] - sub[col_s].values[0]
- else:
- row[f'{zone}_S'] = np.nan
- delta_data.append(row)
- # Create DataFrame
- df_delta_speed = pd.DataFrame(delta_data)
- # Adjust column order
- df_delta_speed = df_delta_speed[
- ['B_id', 'nest_D', 'nest_S', 'edge_D', 'edge_S', 'center_D', 'center_S']
- ]
- # Save results
- df_delta_speed.to_csv(f'data/Fig1G_delta_zones_speed_{period}.csv', index=False)
- # %%
- for period in time_period_order:
- df = pd.read_csv(f'data/Fig1G_delta_zones_speed_{period}.csv')
- groups = ['nest', 'edge', 'center']
- fig, ax = plt.subplots()
- x_pos = {}
- xticks = []
- xticklabels = []
- x = 0
- colors = {'D': 'orange', 'S': 'blue'}
- # 1. x
- for g in groups:
- for cond in ['D', 'S']:
- if cond == 'D':
- x_pos[f'{g}_{cond}_mean'] = x
- x_pos[f'{g}_{cond}_pts'] = x + 0.8
- xticks += [x, x + 0.8]
- xticklabels += [f'{g}_{cond}_mean', f'{g}_{cond}_pts']
- else:
- x_pos[f'{g}_{cond}_pts'] = x
- x_pos[f'{g}_{cond}_mean'] = x + 0.8
- xticks += [x, x + 0.8]
- xticklabels += [f'{g}_{cond}_pts', f'{g}_{cond}_mean']
- x += 2.2
- # 2. mean ± SD
- for g in groups:
- for cond in ['D', 'S']:
- col = f'{g}_{cond}'
- xpos = x_pos[f'{g}_{cond}_mean']
- ax.errorbar(
- xpos,
- df[col].mean(),
- yerr=df[col].std(),
- fmt='o',
- color=colors[cond],
- capsize=4,
- markersize=9,
- elinewidth=1.5,
- capthick=1.5,
- zorder=3
- )
- # 3. raw points + pairing
- for i in range(len(df)):
- for g in groups:
- D_col = f'{g}_D'
- S_col = f'{g}_S'
- x_d = x_pos[f'{g}_D_pts']
- y_d = df.loc[i, D_col]
- x_s = x_pos[f'{g}_S_pts']
- y_s = df.loc[i, S_col]
- ax.scatter(x_d, y_d, color=colors['D'])
- ax.scatter(x_s, y_s, color=colors['S'])
- ax.plot([x_d, x_s], [y_d, y_s],
- color='gray', linewidth=1)
- # 4. axis
- ax.set_xticks(xticks)
- ax.set_xticklabels(xticklabels, rotation=45, ha='right')
- ax.set_ylabel('Δ Speed in zone (cm/s)')
- ax.set_ylim([-40, 40])
- ax.set_title(f'Δ Speed in zone ({period})')
- plt.tight_layout()
- plt.show()
- # %% [markdown]
- # ### Fig1H. post-looming ethogram
- # %%
- import numpy as np
- import pandas as pd
- import matplotlib.pyplot as plt
- import matplotlib.colors as mcolors
- import colorsys
- from matplotlib.patches import FancyArrow
- from analyze_data_utils import filter_in_range
- def adjust_saturation(color, saturation):
- rgb = mcolors.to_rgb(color)
- h, l, s = colorsys.rgb_to_hls(*rgb)
- new_s = saturation * s
- new_rgb = colorsys.hls_to_rgb(h, l, new_s)
- return new_rgb
- with open('data/lst_etho_dict.pkl', 'rb') as f:
- lst_etho_dict = pickle.load(f)
- lst_framerange = [0, 1950]
- saturation = 0.9
- framerate=30
- stim_frame = int(0.725 * framerate)
- behavior_frames_dict = lst_etho_dict
- behavior_frames_dict = dict(sorted(behavior_frames_dict.items()))
- lst_behavior_frames_dict = filter_in_range(lst_etho_dict, lst_framerange, method='replace')
- lst_behavior_properties = {
- 'approach_partner': ('#143FCA', 4),
- 'follow_partner': ('#137CAB', 4),
- 'groom_partner': ('#0990FF', 4),
- 'sniff_partner': ('#09FFFF', 4),
- 'tailrattling': ('#FF09FF', 5),
- 'huddling': ('#6D57F3', 4),
- 'jumping': ('#FF1717', 5),
- 'escape': ('#FF1717', 4),
- 'freezing': ('#FF75EF', 4),
- 'dwelling': ('#FF75EF', 4),
- 'grooming': ('#28AE61', 4),
- 'real_rearing':('#C409FF', 3),
- 'stretching': ('#C409FF', 3),
- 'nest': ('gray', 2),
- 'sniffing': ('#0C8140', 1),
- 'climbing': ('gray', 1),
- }
- fig, axes = plt.subplots(1, 2, figsize=(10, 2), dpi=300)
- plt.subplots_adjust(wspace=0)
- for i, group in enumerate(lst_etho_dict.keys()):
- ax = axes[i]
- ax.set_yticks(range(len(lst_etho_dict[group].keys())))
- for idx, t_id in enumerate(lst_etho_dict[group].keys()):
- behaviors = lst_etho_dict[group][t_id]
- for b, frame_ranges in behaviors.items():
- for frame_range in frame_ranges:
- if np.isnan(frame_ranges).all():
- continue
- start_frame, end_frame = frame_range
- color, zorder = lst_behavior_properties.get(b, ('white', 0))
- color = adjust_saturation(color, saturation)
- rect = plt.Rectangle((start_frame, idx - 0.4), end_frame - start_frame, 0.8, facecolor=color, edgecolor='none', zorder=zorder)
- ax.add_patch(rect)
- if idx % 2 == 0:
- ax.axhline(idx - 0.5, color='black', linestyle='-', linewidth=1, zorder=9)
- ytick_color = 'red'
- else:
- ax.axhline(idx - 0.5, color='black', linestyle='-', linewidth=0.1, zorder=9)
- ytick_color = 'green'
- ax.get_yticklabels()[idx].set_color(ytick_color)
- span_start_frame = 5 * framerate
- span_end_frame = span_start_frame + stim_frame
- if i == 0:
- ax.set_yticklabels(['D', 'S', 'D', 'S', 'D', 'S'])
- ax.set_ylabel('Looming')
- else:
- ax.set_yticks([])
- ax.set_xticks([])
- lst_blank_width = (lst_framerange[1] - lst_framerange[0]) / 20
- per_rect = plt.Rectangle((lst_framerange[0] - lst_blank_width, -0.5), lst_blank_width, 6,
- facecolor='white', edgecolor='white', zorder=8)
- ax.add_patch(per_rect)
- post_rect = plt.Rectangle((lst_framerange[1], -0.5), lst_blank_width, 6,
- facecolor='white', edgecolor='white', zorder=8)
- ax.add_patch(post_rect)
- ax.set_xlim(lst_framerange[0] - lst_blank_width, lst_framerange[1] + lst_blank_width)
- ax.set_ylim(-0.5, 5.5)
- ax.invert_yaxis()
- ax.set_title(f'{group}')
- 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')
- lst_arrow.set_clip_on(False)
- ax.add_patch(lst_arrow)
- threat_line = plt.Line2D([span_start_frame, span_end_frame], [5.75, 5.75], color='red', linewidth=2)
- threat_line.set_clip_on(False)
- ax.add_artist(threat_line)
- if i == 1:
- line_x1 = lst_framerange[1] - 30*10
- line_x2 = lst_framerange[1]
- ax.text((line_x1 + line_x2) / 2, -0.8, '10 sec', fontsize=8, color='black', ha='center')
- lst_line = plt.Line2D([line_x1, line_x2], [-0.65, -0.65], color='black', linewidth=1)
- lst_line.set_clip_on(False)
- ax.add_artist(lst_line)
- lst_legend_ = {
- 'escape': ('#FF1717', 4),
- 'freezing': ('#FF75EF', 4),
- 'tail rattling': ('#FF09FF', 5),
- 'stretching / rearing': ('#C409FF', 3),
- 'huddling': ('#6D57F3', 1),
- 'approaching P (partner)': ('#143FCA', 4),
- 'following P': ('#137CAB', 4),
- 'grooming P': ('#0990FF', 4),
- 'sniffing P': ('#09FFFF', 4),
- 'grooming': ('#28AE61', 4),
- 'sniffing': ('#0C8140', 1),
- 'in nest': ('gray', 2),
- 'other behaviors': ('white', 1)}
- lst_legend = []
- for label, (color, zorder) in lst_legend_.items():
- if label == 'other behaviors':
- rect = plt.Rectangle((0, 0), 1, 1, facecolor=adjust_saturation(color, saturation),
- edgecolor='black', linewidth=0.5, label=label)
- else:
- rect = plt.Rectangle((0, 0), 1, 1, facecolor=adjust_saturation(color, saturation), label=label)
- lst_legend.append(rect)
- axes[1].legend(handles=lst_legend, loc='upper center', bbox_to_anchor=(0, -0.2), fontsize=7, ncol=7)
- plt.savefig('fig/fig_1/lst_behavior_ethogram.eps', format="eps", dpi=300, bbox_inches="tight")
- plt.show()
- # %% [markdown]
- # ### Fig1I. Social time(s/(min∙cm2))
- # %%
- import copy
- import numpy as np
- import pandas as pd
- from analyze_data_utils import dict_to_dataframe, filter_dict_data, merge_dicts, filter_in_range, calculate_duration_time, calculate_total
- def filter_social_in_nest(social_frames, location_series):
- """
- social_frames: [(start, end), ...]
- location_series: pd.Series, starting from 0, 0 indicates in nest
- Returns: [(nest_start, nest_end), ...]
- """
- social_in_nest = []
- for start, end in social_frames:
- # Extract location values for this interval
- loc_segment = location_series[start:end+1] # +1 ensures inclusion of end frame
- # Find consecutive segments equal to 0
- in_nest_mask = (loc_segment == 0)
- if not in_nest_mask.any():
- continue # This social segment is completely outside the nest
- # Find consecutive True segments
- in_nest_indices = loc_segment.index[in_nest_mask]
- group_start = None
- for idx in in_nest_indices:
- if group_start is None:
- group_start = idx
- prev_idx = idx
- elif idx == prev_idx + 1:
- prev_idx = idx
- else:
- social_in_nest.append((group_start, prev_idx))
- group_start = idx
- prev_idx = idx
- # Last segment
- if group_start is not None:
- social_in_nest.append((group_start, prev_idx))
- if social_in_nest == []:
- social_in_nest = [np.nan]
- return social_in_nest
- ap_bl_data = filter_dict_data(lst_baseline_bhvr_labels_dict, 'approach_partner')
- sp_bl_data = filter_dict_data(lst_baseline_bhvr_labels_dict, 'sniff_partner')
- fp_bl_data = filter_dict_data(lst_baseline_bhvr_labels_dict, 'follow_partner')
- gp_bl_data = filter_dict_data(lst_baseline_bhvr_labels_dict, 'groom_partner')
- hp_bl_data = filter_dict_data(lst_baseline_bhvr_labels_dict, 'huddling')
- merged_bl_data = merge_dicts(ap_bl_data, sp_bl_data)
- merged_bl_data = merge_dicts(merged_bl_data, fp_bl_data)
- merged_bl_data = merge_dicts(merged_bl_data, gp_bl_data)
- merged_bl_data = merge_dicts(merged_bl_data, hp_bl_data)
- ap_wr_data = filter_dict_data(lst_bhvr_labels_dict, 'approach_partner')
- sp_wr_data = filter_dict_data(lst_bhvr_labels_dict, 'sniff_partner')
- fp_wr_data = filter_dict_data(lst_bhvr_labels_dict, 'follow_partner')
- gp_wr_data = filter_dict_data(lst_bhvr_labels_dict, 'groom_partner')
- hp_wr_data = filter_dict_data(lst_bhvr_labels_dict, 'huddling')
- merged_wr_data = merge_dicts(ap_wr_data, sp_wr_data)
- merged_wr_data = merge_dicts(merged_wr_data, fp_wr_data)
- merged_wr_data = merge_dicts(merged_wr_data, gp_wr_data)
- merged_wr_data = merge_dicts(merged_wr_data, hp_wr_data)
- merged_bl_data = filter_in_range(merged_bl_data, [0,4500], method='replace')
- # merged_wr_data = filter_in_range(merged_wr_data, [150,1950], method='replace')
- # merged_wr_data = filter_in_range(merged_wr_data, [150,300], method='replace')
- merged_wr_data = filter_in_range(merged_wr_data, [150,1950], method='replace')
- merged_data = {'bl': merged_bl_data['p'], 'pl': merged_wr_data['p']}
- # merged_data_ = merge_intervals(merged_data, max_gap=30)
- merged_data_ = merged_data
- g = 'p'
- social_in_nest_data = {}
- for m in ['D','S']:
- for t in ['bl', 'pl']:
- m_ids = mS_ids if m =='S' else mD_ids
- lst_data = lsts_data if g == 's' else lstp_data
- for m_id in m_ids:
- mp_id = [k for k in mp_ids if m_id in k]
- if t == 'bl':
- interested_frame = (9000, 9000+4500)
- t_id = m_id.split('M')[0]+'M'+m_id.split('(')[1][0]+'T1'
- coord_data = {'x': lst_data[m_id]['coordinate']['waist']['x'][interested_frame[0]:interested_frame[1]].reset_index(drop=True),
- 'y': lst_data[m_id]['coordinate']['waist']['y'][interested_frame[0]:interested_frame[1]].reset_index(drop=True)}
- location_value = get_lst_location_value(coord_data)
- social_frames = merged_data_[t][m][t_id]
- if t not in social_in_nest_data:
- social_in_nest_data[t] = {}
- if m not in social_in_nest_data[t]:
- social_in_nest_data[t][m] = {}
- if not np.isnan(social_frames).any():
- social_in_nest = filter_social_in_nest(social_frames, location_value)
- social_in_nest_data[t][m][t_id] = social_in_nest
- else:
- social_in_nest_data[t][m][t_id] = [np.nan]
- elif t == 'pl':
- for idx, lsf in enumerate(looming_start_frames[mp_id[0]], start=1):
- interested_frame = (lsf+30*0, lsf+30*60)
- t_id = m_id.split('M')[0]+'M'+m_id.split('(')[1][0]+'T'+str(idx)
- # interes20ted_frame = (looming_start_frames[m_id][0], looming_start_frames[m_id][-1]+30*60)
- # for lsf in looming_start_frames[m_id]:
- # t_id = m_id.split('M')[0]+'M'+m_id.split('(')[1][0]+'T1'
- # for nest in nest_data[g][m][t_id]:
- # if not np.isnan(nest).any():
- # start, end = nest
- # interested_frame = (9000+start-60, 9000+start+10)
- # interested_frame = (lsf-150+start-60, lsf-150+start+10)
- coord_data = {'x': lst_data[m_id]['coordinate']['waist']['x'][interested_frame[0]:interested_frame[1]].reset_index(drop=True),
- 'y': lst_data[m_id]['coordinate']['waist']['y'][interested_frame[0]:interested_frame[1]].reset_index(drop=True)}
- location_value = get_lst_location_value(coord_data)
- if t_id in merged_data_[t][m].keys():
- social_frames = merged_data_[t][m][t_id]
- if t not in social_in_nest_data:
- social_in_nest_data[t] = {}
- if m not in social_in_nest_data[t]:
- social_in_nest_data[t][m] = {}
- if not np.isnan(social_frames).any():
- social_in_nest = filter_social_in_nest(social_frames, location_value)
- social_in_nest_data[t][m][t_id] = social_in_nest
- else:
- social_in_nest_data[t][m][t_id] = [np.nan]
- raw_social_in_nest_data = copy.deepcopy(social_in_nest_data)
- social_in_nest_data = {'bl': raw_social_in_nest_data['bl'], 'pl': raw_social_in_nest_data['pl']}
- social_in_nest_durations = calculate_duration_time(social_in_nest_data)
- total_social_in_nest_duration = calculate_total(social_in_nest_durations)
- total_social_in_nest_duration_df = dict_to_dataframe(total_social_in_nest_duration, value_name='in_nest_duration', groups=['bl', 'pl'], nan2zero=True)
- # total_social_in_nest_duration_df.to_csv('data/lst_total_social_in_nest_duration.csv')
- raw_social_data = copy.deepcopy(merged_data_)
- social_data = {'bl': raw_social_data['bl'], 'pl': raw_social_data['pl']}
- social_durations = calculate_duration_time(social_data)
- total_social_duration = calculate_total(social_durations)
- total_social_duration_df = dict_to_dataframe(total_social_duration, value_name='total_social_duration', groups=['bl', 'pl'], nan2zero=True)
- # Create total_social_out_nest_duration_df, calculate difference for each suffixed column separately
- total_social_out_nest_duration_df = total_social_duration_df[['B_id']].copy()
- # Calculate difference for each suffix (_blD, _blS, _plD, _plS) separately
- for suffix in ['_blD', '_blS', '_plD', '_plS']:
- total_col = f'total_social_duration{suffix}'
- in_nest_col = f'in_nest_duration{suffix}'
- out_nest_col = f'out_nest_duration{suffix}'
- if total_col in total_social_duration_df.columns and in_nest_col in total_social_in_nest_duration_df.columns:
- total_social_out_nest_duration_df[out_nest_col] = (
- total_social_duration_df[total_col] - total_social_in_nest_duration_df[in_nest_col]
- )
- # total_social_out_nest_duration_df.to_csv('data/lst_total_social_out_nest_duration.csv')
- # Merge in_nest and out_nest data, in first, out later
- # Define parameters
- baseline_time = 2.5 # 2.5 min = 150 seconds
- afterlooming_time = 1 # 1 min = 60 seconds
- in_nest_area = 300 # cm²
- out_nest_area = 2500 - 300 # 2200 cm²
- # Merge DataFrame
- social_density_df = total_social_in_nest_duration_df[['B_id']].copy()
- # Add columns in order: in_nest first, then out_nest
- for condition in ['_blD', '_blS', '_plD', '_plS']:
- # Add in_nest column first
- in_col = f'in_nest_duration{condition}'
- if in_col in total_social_in_nest_duration_df.columns:
- # Determine if baseline (bl) or after-looming (pl)
- time_divisor = baseline_time if '_bl' in condition else afterlooming_time
- # Calculate density: (duration / time) / area
- social_density_df[in_col] = (
- total_social_in_nest_duration_df[in_col] / time_divisor / in_nest_area
- )
- # Add out_nest column later
- out_col = f'out_nest_duration{condition}'
- if out_col in total_social_out_nest_duration_df.columns:
- # Determine if baseline (bl) or after-looming (pl)
- time_divisor = baseline_time if '_bl' in condition else afterlooming_time
- # Calculate density: (duration / time) / area
- social_density_df[out_col] = (
- total_social_out_nest_duration_df[out_col] / time_divisor / out_nest_area
- )
- social_density_df.to_csv('data/Fig1I_social_time_density.csv', index=False)
- # %%
- df = pd.read_csv('data/Fig1I_social_time_density.csv')
- 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']
- fig, ax = plt.subplots()
- x_pos = {}
- xticks = []
- xticklabels = []
- x = 0
- colors = {'D': 'orange', 'S': 'blue'}
- for col in ordered_cols:
- cond = col[-1]
- is_out = 'out' in col
- if not is_out:
- x_pos[f'{col}_mean'] = x
- x_pos[f'{col}_pts'] = x + 0.8
- xticks += [x, x + 0.8]
- xticklabels += [f'{col}_mean', f'{col}_pts']
- else:
- x_pos[f'{col}_pts'] = x
- x_pos[f'{col}_mean'] = x + 0.8
- xticks += [x, x + 0.8]
- xticklabels += [f'{col}_pts', f'{col}_mean']
- x += 2.2
- for col in ordered_cols:
- cond = col[-1]
- 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)
- for i in range(len(df)):
- for c in ['bl', 'pl']:
- for cond in ['D', 'S']:
- in_col = f'in_nest_duration_{c}{cond}'
- out_col = f'out_nest_duration_{c}{cond}'
- x_in = x_pos[f'{in_col}_pts']
- x_out = x_pos[f'{out_col}_pts']
- y_in = df.loc[i, in_col]
- y_out = df.loc[i, out_col]
- ax.plot([x_in, x_out], [y_in, y_out], color='gray', linewidth=1)
- ax.scatter(x_in, y_in, color=colors[cond])
- ax.scatter(x_out, y_out, color=colors[cond])
- ax.set_xticks(xticks)
- ax.set_xticklabels(xticklabels, rotation=45, ha='right')
- ax.set_ylabel('Social time (s/(min﹡cm²))')
- plt.tight_layout()
- plt.show()
- # %% [markdown]
- # ### Fig1K. Allocation of behavioral decisions (%)
- # %%
- import numpy as np
- import pandas as pd
- combine = True
- row_names = ['E', 'F+E', 'F', 'I']
- colors = ['red', 'purple', 'blue', 'green']
- exclude = ['N', 'n']
- # exclude = ['N', 'n', 'I', 'A']
- def classify_behaviors(behaviors_dict, key1, key2, all_behaviors, exclude=None, combine=False):
- behaviors_classified = []
- for k1 in key1:
- for k2 in key2:
- behaviors_classified.extend(list(behaviors_dict[k1][k2].values()))
- if exclude:
- behaviors_classified = [behavior for behavior in behaviors_classified if all(excl not in behavior for excl in exclude)]
- if combine:
- behaviors_classified = ['E' if behavior == 'A+E' else 'I' if behavior == 'A' else behavior for behavior in behaviors_classified]
- behaviors, indices = np.unique(behaviors_classified, return_inverse=True)
- counts = np.bincount(indices)
- percents = np.bincount(indices) / len(behaviors_classified) * 100
- n = len(behaviors_classified)
- behavior_counts = {behavior: 0 for behavior in all_behaviors}
- behavior_percents = {behavior: 0.0 for behavior in all_behaviors}
- for behavior, count, percent in zip(behaviors, counts, percents):
- behavior_counts[behavior] = count
- behavior_percents[behavior] = percent
- sorted_counts = [behavior_counts[behavior] for behavior in all_behaviors]
- sorted_percents = [behavior_percents[behavior] for behavior in all_behaviors]
- return sorted_percents, sorted_counts, n
- def calculate_behavior_decision_percentages(bhvr_deci_dict, behavior_names, exclude_list=None, combine_flag=False):
- sD_deci_pct, sD_deci_cnt, sD_deci_n = classify_behaviors(
- bhvr_deci_dict, ['s'], ['D'], behavior_names, exclude=exclude_list, combine=combine_flag)
- pD_deci_pct, pD_deci_cnt, pD_deci_n = classify_behaviors(
- bhvr_deci_dict, ['p'], ['D'], behavior_names, exclude=exclude_list, combine=combine_flag)
- sS_deci_pct, sS_deci_cnt, sS_deci_n = classify_behaviors(
- bhvr_deci_dict, ['s'], ['S'], behavior_names, exclude=exclude_list, combine=combine_flag)
- pS_deci_pct, pS_deci_cnt, pS_deci_n = classify_behaviors(
- bhvr_deci_dict, ['p'], ['S'], behavior_names, exclude=exclude_list, combine=combine_flag)
- pct_df = pd.DataFrame({
- 'behavior': behavior_names,
- 'single_Dom': sD_deci_pct,
- 'single_Sub': sS_deci_pct,
- 'pair_Dom': pD_deci_pct,
- 'pair_Sub': pS_deci_pct
- })
- cnt_df = pd.DataFrame({
- 'behavior': behavior_names,
- 'single_Dom': sD_deci_cnt,
- 'single_Sub': sS_deci_cnt,
- 'pair_Dom': pD_deci_cnt,
- 'pair_Sub': pS_deci_cnt
- })
- n_df = pd.DataFrame({
- 'group_rank': ['single_Dom', 'single_Sub', 'pair_Dom', 'pair_Sub'],
- 'n': [sD_deci_n, sS_deci_n, pD_deci_n, pS_deci_n]
- })
- return pct_df, cnt_df, n_df
- pct_df, cnt_df, n_df = calculate_behavior_decision_percentages(
- lst_deci_dict, row_names, exclude_list=exclude, combine_flag=combine)
- print(cnt_df)
- pct_df.to_csv('data/Fig1K_lst_behavior_decision_pct.csv', index=False)
- cnt_df.to_csv('data/Fig1K_lst_behavior_decision_cnt.csv', index=False)
- # %%
- import pandas as pd
- import numpy as np
- import matplotlib.pyplot as plt
- df = pd.read_csv('data/Fig1K_lst_behavior_decision_pct.csv')
- colors = {'E': 'red', 'F+E': 'purple', 'F': 'blue', 'I': 'green'}
- x_labels = ['single_Dom', 'pair_Dom', 'single_Sub', 'pair_Sub']
- x = np.arange(len(x_labels))
- width = 0.6
- fig, ax = plt.subplots()
- bottom = np.zeros(len(x_labels))
- for i, row in df.iloc[::-1].iterrows(): # reverse
- vals = [row['single_Dom'], row['pair_Dom'], row['single_Sub'], row['pair_Sub']]
- c = colors[row['behavior']]
- ax.bar(x, vals, width, bottom=bottom, color=c)
- bottom += vals
- ax.set_xticks(x)
- ax.set_xticklabels(x_labels)
- ax.set_ylabel('Allocation of bahavioral decisitons (%)')
- ax.set_ylim(0, 100)
- plt.tight_layout()
- plt.show()
- # %% [markdown]
- # ### Fig1L. Partner-self behavior combinations
- # %%
- import pandas as pd
- import numpy as np
- # Step 1: Merge all behaviors of a trial into a single categorical series
- def merge_lst_behaviors_to_categories(trial_data, frame_range):
- start_frame, end_frame = frame_range
- categorized = pd.Series(['O'] * (end_frame - start_frame),
- index=range(start_frame, end_frame))
- escape_behaviors = ['escape']
- freezing_behaviors = ['freezing', 'tail_rattling']
- social_behaviors = ['approach_partner', 'sniff_partner', 'follow_partner', 'groom_partner', 'huddling']
- for behavior in social_behaviors:
- if behavior in trial_data:
- behavior_series = trial_data[behavior]
- for frame_idx in categorized.index:
- if frame_idx in behavior_series.index and behavior_series[frame_idx] == 1:
- categorized[frame_idx] = 'S'
- for behavior in freezing_behaviors:
- if behavior in trial_data:
- behavior_series = trial_data[behavior]
- for frame_idx in categorized.index:
- if frame_idx in behavior_series.index and behavior_series[frame_idx] == 1:
- categorized[frame_idx] = 'F'
- for behavior in escape_behaviors:
- if behavior in trial_data:
- behavior_series = trial_data[behavior]
- for frame_idx in categorized.index:
- if frame_idx in behavior_series.index and behavior_series[frame_idx] == 1:
- categorized[frame_idx] = 'E'
- return categorized
- # Step 2: Pair partner and self behavior series to generate combination series
- def create_partner_self_combinations(self_series, partner_series):
- """
- Pair partner and self behavior series to produce combination labels.
- Parameters:
- self_series: pandas Series, self behavior classification
- partner_series: pandas Series, partner behavior classification
- Returns:
- pandas Series with combination labels like 'E&F', 'F&F', 'O&O', etc.
- """
- if len(self_series) != len(partner_series):
- raise ValueError("self_series and partner_series length mismatch")
- combination_series = pd.Series(
- [f"{partner_series.iloc[i]}&{self_series.iloc[i]}"
- for i in range(len(self_series))],
- index=self_series.index
- )
- return combination_series
- # Step 3: Calculate time percentage for each combination
- def calculate_combination_percentages(combination_series, all_combinations):
- """
- Compute percentage of frames for each combination.
- Parameters:
- combination_series: pandas Series with combination labels
- Returns:
- dict with percentages for all combinations
- """
- total_frames = len(combination_series)
- percentages = {}
- for combo in all_combinations:
- count = (combination_series == combo).sum()
- percentages[combo] = (count / total_frames * 100) if total_frames > 0 else 0.0
- return percentages
- # Step 4: Helper function to get partner session ID
- def get_partner_session_id(session_id):
- """
- Derive partner session ID from a given session ID.
- Example: B1MDT1 -> B1MST1, B1MST1 -> B1MDT1
- """
- parts = session_id.split('M')
- if len(parts) != 2:
- return None
- base = parts[0] # e.g. "B1"
- rest = parts[1] # e.g. "DT1" or "ST1"
- if rest.startswith('D'):
- partner_rest = 'S' + rest[1:] # D -> S
- elif rest.startswith('S'):
- partner_rest = 'D' + rest[1:] # S -> D
- else:
- return None
- return base + 'M' + partner_rest
- # Step 5: Main processing – analyse all pair trials
- frame_start = 150
- frame_end = 300
- frame_range = (frame_start, frame_end)
- lst_all_combs = ['E&E', 'E&F', 'E&S', 'E&O', 'F&E', 'F&F', 'F&S', 'F&O',
- 'S&E', 'S&F', 'S&S', 'S&O', 'O&E', 'O&F', 'O&S', 'O&O']
- # Initialize result dictionary
- lst_comb_bhvr_dict = {'p': {'D': {}, 'S': {}}}
- # Iterate over all pair group trials
- group = 'p'
- if group in lst_bhvr_frames_dict:
- for rank in ['D', 'S']:
- if rank not in lst_bhvr_frames_dict[group]:
- continue
- for session_id in sorted(lst_bhvr_frames_dict[group][rank].keys()):
- # Get self behavior data
- self_trial_data = lst_bhvr_frames_dict[group][rank][session_id]
- # Determine partner rank and session ID
- partner_rank = 'S' if rank == 'D' else 'D'
- partner_session_id = get_partner_session_id(session_id)
- if partner_session_id is None:
- continue
- # Check if partner data exists
- if (partner_rank not in lst_bhvr_frames_dict[group] or
- partner_session_id not in lst_bhvr_frames_dict[group][partner_rank]):
- continue
- partner_trial_data = lst_bhvr_frames_dict[group][partner_rank][partner_session_id]
- # Step 1: Categorize self and partner behaviors separately
- self_categorized = merge_lst_behaviors_to_categories(self_trial_data, frame_range)
- partner_categorized = merge_lst_behaviors_to_categories(partner_trial_data, frame_range)
- # Step 2: Create partner‑self combination series
- combination_series = create_partner_self_combinations(
- self_categorized, partner_categorized
- )
- # Step 3: Compute percentages for all 16 combinations
- percentages = calculate_combination_percentages(combination_series, lst_all_combs)
- # Step 4: Store in result dictionary
- lst_comb_bhvr_dict['p'][rank][session_id] = percentages
- # %%
- from analyze_data_utils import filter_dict_data, dict_to_dataframe
- import matplotlib.pyplot as plt
- import seaborn as sns
- import numpy as np
- # Store DataFrames for all combinations
- all_dfs = {}
- # Process each combination
- for comb in lst_all_combs:
- # Filter data for this combination
- lst_comb_bhvr_dict_filtered = filter_dict_data(lst_comb_bhvr_dict, comb)
- # Convert to DataFrame
- lst_comb_bhvr_df = dict_to_dataframe(
- lst_comb_bhvr_dict_filtered,
- groups=['p'],
- ranks=['D', 'S'],
- value_name=f'{comb}_pct',
- nan2zero=True
- )
- all_dfs[comb] = lst_comb_bhvr_df
- # Merge all combination data and compute average for dominant mice
- dom_avg_data = {}
- for comb in lst_all_combs:
- df = all_dfs[comb]
- # Select columns for dominant mice (column names ending with '_pD')
- dom_cols = [col for col in df.columns if col.endswith('_pD')]
- if len(dom_cols) > 0:
- dom_values = df[dom_cols].values.flatten()
- dom_avg_data[comb] = dom_values.mean() if len(dom_values) > 0 else 0.0
- else:
- dom_avg_data[comb] = 0.0
- # Build a 4x4 matrix from dom_avg_data
- categories = ['E', 'F', 'S', 'O']
- matrix_4x4 = np.zeros((4, 4))
- for i, cat1 in enumerate(categories):
- for j, cat2 in enumerate(categories):
- comb = f"{cat1}&{cat2}"
- if comb in dom_avg_data:
- matrix_4x4[i, j] = dom_avg_data[comb]
- # Create 4x4 DataFrame
- lst_dom_bhvr_comb_4x4 = pd.DataFrame(
- matrix_4x4,
- index=categories,
- columns=categories
- )
- # Drop 'O' row and column to get 3x3
- lst_dom_bhvr_comb_3x3 = lst_dom_bhvr_comb_4x4.drop(index='O', columns='O')
- # Normalize 3x3 matrix so that total sums to 100%
- matrix_sum = lst_dom_bhvr_comb_3x3.values.sum()
- if matrix_sum > 0:
- lst_dom_bhvr_comb_3x3_normalized = lst_dom_bhvr_comb_3x3 / matrix_sum * 100
- else:
- lst_dom_bhvr_comb_3x3_normalized = lst_dom_bhvr_comb_3x3
- # Save normalized 3x3 matrix
- lst_dom_bhvr_comb_3x3_normalized.to_csv('data/Fig1L_dom_bhvr_comb.csv')
- # %%
- lst_dom_bhvr_comb_3x3_normalized = pd.read_csv('data/Fig1L_dom_bhvr_comb.csv', index_col=0)
- plt.figure()
- sns.heatmap(lst_dom_bhvr_comb_3x3_normalized, annot=True, fmt='.2f', cmap='viridis', vmin=0)
- plt.gca().invert_yaxis()
- plt.xlabel('Dom Behavior')
- plt.ylabel('Sub Behavior')
- plt.show()
- # %%
- from analyze_data_utils import analyze_independence
- # Compute actual total_observations from original data
- # Get number of dominant trials
- num_dom_trials = len(lst_bhvr_frames_dict['p']['D'])
- frames_per_trial = frame_end - frame_start # 300 - 150 = 150 frames
- total_observations = num_dom_trials * frames_per_trial # total frames
- print(f"Number of dominant trials: {num_dom_trials}")
- print(f"Frames per trial: {frames_per_trial}")
- print(f"Total observations (total_observations): {total_observations}")
- # Convert percentages to counts
- # lst_dom_bhvr_comb_3x3_normalized contains normalized percentages (sum=100%)
- # Convert to actual frame counts
- sub_dom_raw_behaviors = (lst_dom_bhvr_comb_3x3_normalized.values / 100 * total_observations).round().astype(int)
- behavior_labels = lst_dom_bhvr_comb_3x3_normalized.index.tolist()
- # Run full independence analysis
- results_behaviors = analyze_independence(
- sub_dom_raw=sub_dom_raw_behaviors,
- labels=behavior_labels,
- title="LST: Behavior Combinations Independence Test (Defense Time)",
- data_type="time"
- )
- # %% [markdown]
- # ### Fig1M. Defense time (%)
- # %%
- import numpy as np
- import pandas as pd
- from analyze_data_utils import calculate_total, filter_dict_data, filter_in_range, calculate_duration_time, dict_to_dataframe, merge_dicts
- # Define framerate and time periods
- framerate = 30
- time_periods = {
- 'baseline': {'range': [0, 4500], 'total_time': 150}, # 150s
- '0-5s': {'range': [150, 300], 'total_time': 5}, # 5s
- '5-20s': {'range': [300, 750], 'total_time': 15}, # 15s
- '20-60s': {'range': [750, 1950], 'total_time': 40} # 40s
- }
- # Extract defense data (escape, freezing, rearing, tail-rattling) for baseline and looming periods
- escape_data_bl = filter_dict_data(lst_baseline_bhvr_labels_dict, 'escape')
- freezing_data_bl = filter_dict_data(lst_baseline_bhvr_labels_dict, 'freezing')
- rearing_data_bl = filter_dict_data(lst_baseline_bhvr_labels_dict, 'rearing/up_stretch')
- tailrattling_data_bl = filter_dict_data(lst_baseline_bhvr_labels_dict, 'tail_rattling')
- escape_data = filter_dict_data(lst_bhvr_labels_dict, 'escape')
- freezing_data = filter_dict_data(lst_bhvr_labels_dict, 'freezing')
- rearing_data = filter_dict_data(lst_bhvr_labels_dict, 'rearing/up_stretch')
- tailrattling_data = filter_dict_data(lst_bhvr_labels_dict, 'tail_rattling')
- # Merge all defense behaviors
- defense_data_bl = merge_dicts(escape_data_bl, freezing_data_bl)
- defense_data_bl = merge_dicts(defense_data_bl, rearing_data_bl)
- defense_data_bl = merge_dicts(defense_data_bl, tailrattling_data_bl)
- defense_data = merge_dicts(escape_data, freezing_data)
- defense_data = merge_dicts(defense_data, rearing_data)
- defense_data = merge_dicts(defense_data, tailrattling_data)
- # Store percentage data for all time periods
- all_percentage_data = {}
- # Process baseline period - use replace method to clip frames within range
- defense_data_bl_filtered = filter_in_range(defense_data_bl, time_periods['baseline']['range'], method='replace')
- duration_bl = calculate_duration_time(defense_data_bl_filtered, framerate=framerate)
- duration_bl = calculate_total(duration_bl)
- duration_bl_df = dict_to_dataframe(duration_bl, value_name='duration', nan2zero=True)
- all_percentage_data['baseline'] = duration_bl_df
- # Process 3 looming periods - use replace method to clip frames within range
- for period_name in ['0-5s', '5-20s', '20-60s']:
- period_range = time_periods[period_name]['range']
- defense_period = filter_in_range(defense_data, period_range, method='replace')
- duration_period = calculate_duration_time(defense_period, framerate=framerate)
- duration_period = calculate_total(duration_period)
- duration_period_df = dict_to_dataframe(duration_period, value_name='duration', nan2zero=True)
- all_percentage_data[period_name] = duration_period_df
- groups = ['s', 'p']
- ranks = ['D', 'S']
- # Store final data
- final_data = {}
- for period_name, df in all_percentage_data.items():
- total_time = time_periods[period_name]['total_time']
- row_data = []
- # Extract data in the order: sD, sS, pD, pS
- for group in groups:
- for rank in ranks:
- col_name = f'duration_{group}{rank}'
- if col_name in df.columns:
- # Calculate percentages and keep original order (sorted by B_id)
- percentages = (df[col_name] / total_time) * 100
- row_data.extend(percentages.tolist())
- final_data[period_name] = row_data
- # Create final DataFrame
- final_df = pd.DataFrame(final_data).T
- # Generate column names: sD_1, sD_2, ..., sS_1, sS_2, ..., pD_1, pD_2, ..., pS_1, pS_2, ...
- column_names = []
- first_df = list(all_percentage_data.values())[0]
- for group in groups:
- for rank in ranks:
- col_name = f'duration_{group}{rank}'
- if col_name in first_df.columns:
- count = len(first_df)
- column_names.extend([f'{group}{rank}_{i+1}' for i in range(count)])
- final_df.columns = column_names
- # Final result
- defense_percentage_df = final_df
- # Save result
- defense_percentage_df.to_csv('data/Fig1M_defense_time_pct.csv')
- # %%
- import pandas as pd
- import numpy as np
- import matplotlib.pyplot as plt
- df = pd.read_csv('data/Fig1M_defense_time_pct.csv')
- time = df['Unnamed: 0'].values
- groups = ['sD', 'sS', 'pD', 'pS']
- fig, ax = plt.subplots()
- style = {
- 'sD': {'color': 'orange', 'linestyle': '--', 'marker': 'o', 'mfc': 'none'},
- 'sS': {'color': 'blue', 'linestyle': '--', 'marker': 'o', 'mfc': 'none'},
- 'pD': {'color': 'orange', 'linestyle': '-', 'marker': 'o', 'mfc': 'orange'},
- 'pS': {'color': 'blue', 'linestyle': '-', 'marker': 'o', 'mfc': 'blue'}
- }
- for g in groups:
- cols = [c for c in df.columns if c.startswith(g + '_')]
- means = df[cols].mean(axis=1)
- stds = df[cols].sem(axis=1)
- ax.errorbar(
- time,
- means,
- yerr=stds,
- color=style[g]['color'],
- linestyle=style[g]['linestyle'],
- marker=style[g]['marker'],
- markerfacecolor=style[g]['mfc'],
- capsize=4,
- linewidth=2,
- label=g
- )
- ax.set_xticks(range(len(time)))
- ax.set_xticklabels(time, rotation=45, ha='right')
- ax.set_ylabel('Percentage')
- ax.set_ylim(0, 60)
- ax.legend()
- plt.tight_layout()
- plt.show()
- # %% [markdown]
- # ### Fig1N. Peak escape speed (cm/s)
- # %%
- import numpy as np
- from analyze_data_utils import filter_dict_data, filter_in_range, calculate_latency_time, dict_to_dataframe
- escape_data = filter_dict_data(lst_bhvr_labels_dict, 'escape')
- escape_data = filter_in_range(escape_data, [150,1950], method='delete')
- escape_latency = calculate_latency_time(escape_data)
- escape_latency_df = dict_to_dataframe(escape_latency, value_name='escape_latency')
- escape_latency_df.to_csv('data/lst_1st_escape_latency.csv', index=False)
- vel_escape_data = {'p':{'D': {}, 'S': {}}, 's': {'D': {}, 'S': {}}}
- vel_escape_latency = {'p':{'D': {}, 'S': {}}, 's': {'D': {}, 'S': {}}}
- for g in ['s', 'p']:
- for r in ['D', 'S']:
- for t_id in escape_data[g][r].keys():
- if t_id.split('B')[1].split('M')[0] in ['1', '3', '7', '10', '11', '12', '13', '14', '15', '16', '17', '20']:
- m = 1 if t_id.split('M')[1].split('T')[0] == 'S' else 2
- else:
- m = 1 if t_id.split('M')[1].split('T')[0] == 'D' else 2
- m_id = t_id.split('M')[0] + 'M' + str(m) + '(' + t_id.split('M')[1].split('T')[0] + ')'
- if g == 's':
- velocity_data = lsts_data[m_id]['velocity']['waist']
- looming_start_frame = looming_start_frames[m_id][int(t_id.split('T')[1])-1]
- elif g == 'p':
- velocity_data = lstp_data[m_id]['velocity']['waist']
- mp_id = mp_ids[int(m_id.split('B')[1].split('M')[0])-1]
- looming_start_frame = looming_start_frames[mp_id][int(t_id.split('T')[1])-1]
- escape_list = escape_data[g][r][t_id]
- for escape in escape_list:
- if np.any(np.isnan(escape)):
- continue
- else:
- start, end = escape
- escape_start_frame = looming_start_frame+start-150
- escape_end_frame = looming_start_frame+end-150
- vel_nest = velocity_data[escape_start_frame:escape_end_frame].max() / lst_pixpercm * framerate
- if t_id not in vel_escape_data[g][r].keys():
- vel_escape_data[g][r][t_id] = []
- vel_escape_data[g][r][t_id].append(vel_nest)
- max_velocity_index = np.argmax(velocity_data[escape_start_frame:escape_end_frame])
- max_velocity_frame = escape_start_frame + max_velocity_index
- latency_to_max_velocity = (max_velocity_frame - escape_start_frame) / framerate
- if t_id in ['B10MDT3', 'B10MDT4', 'B10MDT5', 'B10MDT9']:
- # Get the velocity segment for these specific trials
- velocity_segment = velocity_data[escape_start_frame:escape_end_frame].values
- velocity_segment_cm_s = velocity_segment / lst_pixpercm * framerate
- if t_id not in vel_escape_latency[g][r].keys():
- vel_escape_latency[g][r][t_id] = []
- vel_escape_latency[g][r][t_id].append(latency_to_max_velocity)
- vel_escape_df = dict_to_dataframe(vel_escape_data, value_name='escape_speed')
- vel_escape_df.to_csv('data/Fig1N_max_escape_speed.csv', index=False)
- # %%
- import pandas as pd
- import numpy as np
- import matplotlib.pyplot as plt
- df = pd.read_csv('data/Fig1N_max_escape_speed.csv')
- groups = ['escape_speed_sD', 'escape_speed_pD', 'escape_speed_sS', 'escape_speed_pS']
- fig, ax = plt.subplots()
- x_pos = {}
- xticks = []
- xticklabels = []
- x = 0
- colors = {'sD': 'orange', 'sS': 'blue', 'pD': 'orange', 'pS': 'blue'}
- # 1. x
- for g in groups:
- cond = g.split('_')[-1]
- if cond in ['sD', 'sS']:
- x_pos[f'{g}_mean'] = x
- x_pos[f'{g}_pts'] = x + 0.8
- xticks += [x, x + 0.8]
- xticklabels += [f'{g}_mean', f'{g}_pts']
- else:
- x_pos[f'{g}_pts'] = x
- x_pos[f'{g}_mean'] = x + 0.8
- xticks += [x, x + 0.8]
- xticklabels += [f'{g}_pts', f'{g}_mean']
- x += 2.2
- # 2. mean ± SD
- for g in groups:
- cond = g.split('_')[-1]
- ax.errorbar(
- x_pos[f'{g}_mean'],
- df[g].mean(),
- yerr=df[g].std(),
- fmt='o',
- color=colors[cond],
- capsize=4,
- markersize=9,
- elinewidth=1.5,
- capthick=1.5,
- zorder=3
- )
- # 3. raw + correct pairing
- for i in range(len(df)):
- sD_x = x_pos['escape_speed_sD_pts']
- sS_x = x_pos['escape_speed_sS_pts']
- pD_x = x_pos['escape_speed_pD_pts']
- pS_x = x_pos['escape_speed_pS_pts']
- sD_y = df.loc[i, 'escape_speed_sD']
- sS_y = df.loc[i, 'escape_speed_sS']
- pD_y = df.loc[i, 'escape_speed_pD']
- pS_y = df.loc[i, 'escape_speed_pS']
- ax.scatter(sD_x, sD_y, color=colors['sD'])
- ax.scatter(sS_x, sS_y, color=colors['sS'])
- ax.scatter(pD_x, pD_y, color=colors['pD'])
- ax.scatter(pS_x, pS_y, color=colors['pS'])
- ax.plot([sD_x, pD_x], [sD_y, pD_y], color='gray', linewidth=1)
- ax.plot([sS_x, pS_x], [sS_y, pS_y], color='gray', linewidth=1)
- # 4. axis
- ax.set_xticks(xticks)
- ax.set_xticklabels(xticklabels, rotation=45, ha='right')
- ax.set_ylabel('Peak escape speed (cm/s)')
- ax.set_ylim(0, 125)
- plt.tight_layout()
- plt.show()
- # %% [markdown]
- # ### Fig1O. First freezing duration (s)
- # %%
- from analyze_data_utils import filter_dict_data, filter_duration, calculate_duration_time
- def get_behavior_frame_when_main_behavior(lst_deci_dict, lst_bhvr_labels_dict, main_behavior, sub_behavior=None):
- target_behavior_data = {'p':{'D': {}, 'S': {}}, 's': {'D': {}, 'S': {}}}
- for g in lst_bhvr_labels_dict.keys():
- for m in lst_bhvr_labels_dict[g].keys():
- for t in lst_bhvr_labels_dict[g][m].keys():
- if lst_deci_dict[g][m][t] in main_behavior:
- if sub_behavior:
- sub_behavior_data = lst_bhvr_labels_dict[g][m][t][sub_behavior]
- else:
- sub_behavior_data = lst_bhvr_labels_dict[g][m][t]
- target_behavior_data[g][m][t] = sub_behavior_data
- return target_behavior_data
- def calculate_bhvr1_duration_before_bhvr2(bhvr1, bhvr2, first=False):
- duration_data = {}
- for key, values in bhvr1.items():
- if isinstance(values, dict):
- duration_data[key] = calculate_bhvr1_duration_before_bhvr2(bhvr1[key], bhvr2[key], first=first)
- else:
- turples = []
- if not np.isnan(values).any():
- for value in values:
- start, end = value
- if np.isnan(bhvr2[key]).any():
- turples.append((start, end))
- else:
- if end <= bhvr2[key][0][0]+1: # 使用第一次escape的开始帧
- turples.append((start, end))
- # 如果需要第一次且turples不为空,只保留第一个
- if first and len(turples) > 0:
- turples = [turples[0]]
- durations = []
- if np.isnan(turples).any():
- duration = np.nan
- durations.append(duration)
- duration_data[key] = durations
- else:
- for turple in turples:
- start, end = turple
- duration = (end - start + 1) / framerate # +1 for inclusive counting
- durations.append(duration)
- duration_data[key] = durations
- return duration_data
- freezing_data = get_behavior_frame_when_main_behavior(lst_deci_dict, lst_bhvr_labels_dict, ['F', 'F+E'], 'freezing')
- # freezing_data = filter_dict_data(behavior_frames_dict, 'freezing')
- escape_data = filter_dict_data(lst_bhvr_labels_dict, 'escape')
- freezing_data = filter_duration(freezing_data)
- freezing_data = filter_in_range(freezing_data, [150, 300], method='replace')
- first_freezing_duration = calculate_bhvr1_duration_before_bhvr2(freezing_data, escape_data, first=True)
- first_freezing_duration_df = dict_to_dataframe(first_freezing_duration, value_name='first_freezing_duration')
- first_freezing_duration_df.to_csv('data/Fig1O_first_freezing_duration.csv', index=False)
- # %%
- import pandas as pd
- import numpy as np
- import matplotlib.pyplot as plt
- df = pd.read_csv('data/Fig1O_first_freezing_duration.csv')
- print(df.columns.tolist())
- groups = ['first_freezing_duration_sD', 'first_freezing_duration_pD', 'first_freezing_duration_sS', 'first_freezing_duration_pS']
- fig, ax = plt.subplots()
- x_pos = {}
- xticks = []
- xticklabels = []
- x = 0
- colors = {'sD': 'orange', 'sS': 'blue', 'pD': 'orange', 'pS': 'blue'}
- # 1. x
- for g in groups:
- cond = g.split('_')[-1]
- if cond in ['sD', 'sS']:
- x_pos[f'{g}_mean'] = x
- x_pos[f'{g}_pts'] = x + 0.8
- xticks += [x, x + 0.8]
- xticklabels += [f'{g}_mean', f'{g}_pts']
- else:
- x_pos[f'{g}_pts'] = x
- x_pos[f'{g}_mean'] = x + 0.8
- xticks += [x, x + 0.8]
- xticklabels += [f'{g}_pts', f'{g}_mean']
- x += 2.2
- # 2. mean ± SD
- for g in groups:
- cond = g.split('_')[-1]
- ax.errorbar(
- x_pos[f'{g}_mean'],
- df[g].mean(),
- yerr=df[g].std(),
- fmt='o',
- color=colors[cond],
- capsize=4,
- markersize=9,
- elinewidth=1.5,
- capthick=1.5,
- zorder=3
- )
- # 3. raw + correct pairing
- for i in range(len(df)):
- sD_x = x_pos['first_freezing_duration_sD_pts']
- sS_x = x_pos['first_freezing_duration_sS_pts']
- pD_x = x_pos['first_freezing_duration_pD_pts']
- pS_x = x_pos['first_freezing_duration_pS_pts']
- sD_y = df.loc[i, 'first_freezing_duration_sD']
- sS_y = df.loc[i, 'first_freezing_duration_sS']
- pD_y = df.loc[i, 'first_freezing_duration_pD']
- pS_y = df.loc[i, 'first_freezing_duration_pS']
- ax.scatter(sD_x, sD_y, color=colors['sD'])
- ax.scatter(sS_x, sS_y, color=colors['sS'])
- ax.scatter(pD_x, pD_y, color=colors['pD'])
- ax.scatter(pS_x, pS_y, color=colors['pS'])
- ax.plot([sD_x, pD_x], [sD_y, pD_y], color='gray', linewidth=1)
- ax.plot([sS_x, pS_x], [sS_y, pS_y], color='gray', linewidth=1)
- # 4. axis
- ax.set_xticks(xticks)
- ax.set_xticklabels(xticklabels, rotation=45, ha='right')
- ax.set_ylabel('1st freezing duration (s)')
- ax.set_ylim(0, 4)
- plt.tight_layout()
- plt.show()
- # %% [markdown]
- # ### Fig1P. Grooming time (%)
- # %%
- import numpy as np
- import pandas as pd
- from analyze_data_utils import calculate_total, filter_dict_data, filter_in_range, calculate_duration_time, dict_to_dataframe
- # Define framerate and time periods
- framerate = 30
- time_periods = {
- 'baseline': {'range': [0, 4500], 'total_time': 150}, # 150s
- '0-5s': {'range': [150, 300], 'total_time': 5}, # 5s
- '5-20s': {'range': [300, 750], 'total_time': 15}, # 15s
- '20-60s': {'range': [750, 1950], 'total_time': 40} # 40s
- }
- # Extract grooming data for baseline and looming periods
- grooming_data_bl = filter_dict_data(lst_baseline_bhvr_labels_dict, 'grooming')
- grooming_data = filter_dict_data(lst_bhvr_labels_dict, 'grooming')
- # Store percentage data for all time periods
- all_percentage_data = {}
- # Process baseline period - use replace method to clip frames within range
- grooming_data_bl_filtered = filter_in_range(grooming_data_bl, time_periods['baseline']['range'], method='replace')
- duration_bl = calculate_duration_time(grooming_data_bl_filtered, framerate=framerate)
- duration_bl = calculate_total(duration_bl)
- duration_bl_df = dict_to_dataframe(duration_bl, value_name='duration', nan2zero=True)
- all_percentage_data['baseline'] = duration_bl_df
- # Process 3 looming periods - use replace method to clip frames within range
- for period_name in ['0-5s', '5-20s', '20-60s']:
- period_range = time_periods[period_name]['range']
- grooming_period = filter_in_range(grooming_data, period_range, method='replace')
- duration_period = calculate_duration_time(grooming_period, framerate=framerate)
- duration_period = calculate_total(duration_period)
- duration_period_df = dict_to_dataframe(duration_period, value_name='duration', nan2zero=True)
- all_percentage_data[period_name] = duration_period_df
- groups = ['s', 'p']
- ranks = ['D', 'S']
- # Store final data
- final_data = {}
- for period_name, df in all_percentage_data.items():
- total_time = time_periods[period_name]['total_time']
- row_data = []
- # Extract data in the order: sD, sS, pD, pS
- for group in groups:
- for rank in ranks:
- col_name = f'duration_{group}{rank}'
- if col_name in df.columns:
- # Calculate percentages and keep original order (sorted by B_id)
- percentages = (df[col_name] / total_time) * 100
- row_data.extend(percentages.tolist())
- final_data[period_name] = row_data
- # Create final DataFrame
- final_df = pd.DataFrame(final_data).T
- # Generate column names: sD_1, sD_2, ..., sS_1, sS_2, ..., pD_1, pD_2, ..., pS_1, pS_2, ...
- column_names = []
- first_df = list(all_percentage_data.values())[0]
- for group in groups:
- for rank in ranks:
- col_name = f'duration_{group}{rank}'
- if col_name in first_df.columns:
- count = len(first_df)
- column_names.extend([f'{group}{rank}_{i+1}' for i in range(count)])
- final_df.columns = column_names
- # Final result
- grooming_percentage_df = final_df
- # Save result
- grooming_percentage_df.to_csv('data/Fig1P_grooming_time_pct.csv')
- # %%
- import pandas as pd
- import numpy as np
- import matplotlib.pyplot as plt
- df = pd.read_csv('data/Fig1P_grooming_time_pct.csv')
- time = df['Unnamed: 0'].values
- groups = ['sD', 'sS', 'pD', 'pS']
- fig, ax = plt.subplots()
- style = {
- 'sD': {'color': 'orange', 'linestyle': '--', 'marker': 'o', 'mfc': 'none'},
- 'sS': {'color': 'blue', 'linestyle': '--', 'marker': 'o', 'mfc': 'none'},
- 'pD': {'color': 'orange', 'linestyle': '-', 'marker': 'o', 'mfc': 'orange'},
- 'pS': {'color': 'blue', 'linestyle': '-', 'marker': 'o', 'mfc': 'blue'}
- }
- for g in groups:
- cols = [c for c in df.columns if c.startswith(g + '_')]
- means = df[cols].mean(axis=1)
- stds = df[cols].sem(axis=1)
- ax.errorbar(
- time,
- means,
- yerr=stds,
- color=style[g]['color'],
- linestyle=style[g]['linestyle'],
- marker=style[g]['marker'],
- markerfacecolor=style[g]['mfc'],
- capsize=4,
- linewidth=2,
- label=g
- )
- ax.set_xticks(range(len(time)))
- ax.set_xticklabels(time, rotation=45, ha='right')
- ax.set_ylabel('Percentage')
- ax.set_ylim(0, 20)
- ax.legend()
- plt.tight_layout()
- plt.show()
- # %% [markdown]
- # ### Fig1Q. Rearing & up-stretch time (%)
- # %%
- import numpy as np
- import pandas as pd
- from analyze_data_utils import calculate_total, filter_dict_data, filter_in_range, calculate_duration_time, dict_to_dataframe
- # Define framerate and time periods
- framerate = 30
- time_periods = {
- 'baseline': {'range': [0, 4500], 'total_time': 150}, # 150s
- '0-5s': {'range': [150, 300], 'total_time': 5}, # 5s
- '5-20s': {'range': [300, 750], 'total_time': 15}, # 15s
- '20-60s': {'range': [750, 1950], 'total_time': 40} # 40s
- }
- # Extract rearing/up_stretch data for baseline and looming periods
- rearing_data_bl = filter_dict_data(lst_baseline_bhvr_labels_dict, 'rearing/up_stretch')
- rearing_data = filter_dict_data(lst_bhvr_labels_dict, 'rearing/up_stretch')
- # Store percentage data for all time periods
- all_percentage_data = {}
- # Process baseline period - use replace method to clip frames within range
- rearing_data_bl_filtered = filter_in_range(rearing_data_bl, time_periods['baseline']['range'], method='replace')
- duration_bl = calculate_duration_time(rearing_data_bl_filtered, framerate=framerate)
- duration_bl = calculate_total(duration_bl)
- duration_bl_df = dict_to_dataframe(duration_bl, value_name='duration', nan2zero=True)
- all_percentage_data['baseline'] = duration_bl_df
- # Process 3 looming periods - use replace method to clip frames within range
- for period_name in ['0-5s', '5-20s', '20-60s']:
- period_range = time_periods[period_name]['range']
- rearing_period = filter_in_range(rearing_data, period_range, method='replace')
- duration_period = calculate_duration_time(rearing_period, framerate=framerate)
- duration_period = calculate_total(duration_period)
- duration_period_df = dict_to_dataframe(duration_period, value_name='duration', nan2zero=True)
- all_percentage_data[period_name] = duration_period_df
- groups = ['s', 'p']
- ranks = ['D', 'S']
- # Store final data
- final_data = {}
- for period_name, df in all_percentage_data.items():
- total_time = time_periods[period_name]['total_time']
- row_data = []
- # Extract data in the order: sD, sS, pD, pS
- for group in groups:
- for rank in ranks:
- col_name = f'duration_{group}{rank}'
- if col_name in df.columns:
- # Calculate percentages and keep original order (sorted by B_id)
- percentages = (df[col_name] / total_time) * 100
- row_data.extend(percentages.tolist())
- final_data[period_name] = row_data
- # Create final DataFrame
- final_df = pd.DataFrame(final_data).T
- # Generate column names: sD_1, sD_2, ..., sS_1, sS_2, ..., pD_1, pD_2, ..., pS_1, pS_2, ...
- column_names = []
- first_df = list(all_percentage_data.values())[0]
- for group in groups:
- for rank in ranks:
- col_name = f'duration_{group}{rank}'
- if col_name in first_df.columns:
- count = len(first_df)
- column_names.extend([f'{group}{rank}_{i+1}' for i in range(count)])
- final_df.columns = column_names
- # Save result
- rearing_percentage_df = final_df
- rearing_percentage_df.to_csv('data/Fig1Q_rearing_up_stretch_time_pct.csv')
- # %%
- import pandas as pd
- import numpy as np
- import matplotlib.pyplot as plt
- df = pd.read_csv('data/Fig1Q_rearing_up_stretch_time_pct.csv')
- time = df['Unnamed: 0'].values
- groups = ['sD', 'sS', 'pD', 'pS']
- fig, ax = plt.subplots()
- style = {
- 'sD': {'color': 'orange', 'linestyle': '--', 'marker': 'o', 'mfc': 'none'},
- 'sS': {'color': 'blue', 'linestyle': '--', 'marker': 'o', 'mfc': 'none'},
- 'pD': {'color': 'orange', 'linestyle': '-', 'marker': 'o', 'mfc': 'orange'},
- 'pS': {'color': 'blue', 'linestyle': '-', 'marker': 'o', 'mfc': 'blue'}
- }
- for g in groups:
- cols = [c for c in df.columns if c.startswith(g + '_')]
- means = df[cols].mean(axis=1)
- stds = df[cols].sem(axis=1)
- ax.errorbar(
- time,
- means,
- yerr=stds,
- color=style[g]['color'],
- linestyle=style[g]['linestyle'],
- marker=style[g]['marker'],
- markerfacecolor=style[g]['mfc'],
- capsize=4,
- linewidth=2,
- label=g
- )
- ax.set_xticks(range(len(time)))
- ax.set_xticklabels(time, rotation=45, ha='right')
- ax.set_ylabel('Percentage')
- ax.set_ylim(0, 20)
- ax.legend()
- plt.tight_layout()
- plt.show()
- # %% [markdown]
- # ### FigS1A. Time in zone when different edge_width
- # %%
- import matplotlib.pyplot as plt
- def count_sequence(series):
- count = 0
- target_structures = [
- [1, 1, 1, 1, 1, 2, 2, 2, 2, 2],
- [2, 2, 2, 2, 2, 1, 1, 1, 1, 1],
- ]
- current_structure = []
- approach_index = []
- for i, value in enumerate(series):
- current_structure.append(value)
- if len(current_structure) > 10:
- current_structure = current_structure[1:]
- if current_structure in target_structures:
- count += 1
- if value == 1:
- approach_index.append(i)
- return count, approach_index
- center_pcts_list = []
- edge_pcts_list = []
- nest_pcts_list = []
- trans_bouts_list = []
- for i in range(0, 26):
- edge_width = 43.52*i/2
- center_pcts = []
- edge_pcts = []
- nest_pcts = []
- trans_bouts = []
- for m_id in ms_ids:
- coord_data = {
- 'x': pd.to_numeric(lsts_data[m_id]['coordinate']['centroid']['x'][0:18000], errors='coerce'),
- 'y': pd.to_numeric(lsts_data[m_id]['coordinate']['centroid']['y'][0:18000], errors='coerce')}
- coord_df = pd.DataFrame(coord_data).dropna()
- local_data = get_lst_location_value(coord_data, edge_width)
- center_pcts.append(local_data.value_counts().get(2, 0) / len(local_data) * 100)
- edge_pcts.append(local_data.value_counts().get(1, 0) / len(local_data) * 100)
- nest_pcts.append(local_data.value_counts().get(0, 0) / len(local_data) * 100)
- trans_bout, trans_idx = count_sequence(local_data)
- trans_bouts.append(trans_bout)
- center_pcts_list.append(np.mean(center_pcts))
- edge_pcts_list.append(np.mean(edge_pcts))
- nest_pcts_list.append(np.mean(nest_pcts))
- trans_bouts_list.append(np.mean(trans_bouts))
- # Calculate edge widths in centimeters for each percentage data point
- edge_widths_cm = [43.52 * i / 2 for i in range(len(edge_pcts_list))]
- df_edge_analysis = pd.DataFrame({
- 'edge_width_cm': edge_widths_cm,
- 'edge_pct': edge_pcts_list,
- 'trans_bouts': trans_bouts_list
- })
- df_edge_analysis.to_csv('data/FigS1A_edge_zone_width.csv', index=False)
- # %%
- from kneed import KneeLocator
- df = pd.read_csv('data/FigS1A_edge_zone_width.csv')
- edge_pcts_list = df['edge_pct'].values
- trans_bouts_list = df['trans_bouts'].values
- # Find the knee point using the Kneedle algorithm
- edge_widths_kneed = np.arange(len(edge_pcts_list))
- edge_pcts_kneed = np.array(edge_pcts_list)
- kl = KneeLocator(
- edge_widths_kneed, edge_pcts_kneed,
- curve='concave', direction='increasing',
- online=True
- )
- x0, x1 = edge_widths_kneed.min(), edge_widths_kneed.max()
- x_diff_mapped = kl.x_difference * (x1 - x0) + x0
- y_diff = kl.y_difference
- fig, ax1 = plt.subplots(figsize=(12, 6), dpi=300)
- ax1.plot(edge_pcts_list, color='red', label='Time in Edge Zone')
- # ax1.plot(np.gradient(edge_pcts_list), color='orange', label='ΔTime in Edge Zone')
- # ax1.plot(np.gradient(np.gradient(edge_pcts_list)), color='yellow', label='Δ²Time in Edge Zone')
- ax1.axvline(x=kl.knee, color='green', linestyle='--', linewidth=1.5, label='Knee')
- ax1.plot(x_diff_mapped, y_diff * np.max(edge_pcts_list),
- color='limegreen', linestyle='-.',
- label='Kneedle Difference')
- ax1.set_ylabel('Time in zone (%)')
- ax1.set_xlabel('Edge Width (cm)')
- ax1.grid(True)
- ax1.legend(loc='upper left')
- ax2 = ax1.twinx()
- ax2.plot(trans_bouts_list, color='darkgray', label='Edge <=> Center Transition Bouts')
- ax2.set_ylabel('Transition Bouts')
- ax2.legend(loc='upper right')
- plt.tight_layout()
- plt.show()
- # %% [markdown]
- # ### FigS1C. baseline ethogram
- # %%
- import pandas as pd
- import matplotlib.pyplot as plt
- import matplotlib.colors as mcolors
- import colorsys
- from matplotlib.patches import FancyArrow
- from analyze_data_utils import filter_in_range
- def adjust_saturation(color, saturation):
- rgb = mcolors.to_rgb(color)
- h, l, s = colorsys.rgb_to_hls(*rgb)
- new_s = saturation * s
- new_rgb = colorsys.hls_to_rgb(h, l, new_s)
- return new_rgb
- with open('data/lst_bl_etho_dict.pkl', 'rb') as f:
- lst_bl_etho_dict = pickle.load(f)
- lst_framerange = [0, 4500]
- saturation = 0.9
- framerate=30
- stim_frame = int(0.725 * framerate)
- behavior_frames_dict = lst_etho_dict
- behavior_frames_dict = dict(sorted(behavior_frames_dict.items()))
- lst_behavior_frames_dict = filter_in_range(lst_etho_dict, lst_framerange, method='replace')
- lst_behavior_properties = {
- 'approach_partner': ('#143FCA', 4),
- 'follow_partner': ('#137CAB', 4),
- 'groom_partner': ('#0990FF', 4),
- 'sniff_partner': ('#09FFFF', 4),
- 'tailrattling': ('#FF09FF', 5),
- 'huddling': ('#6D57F3', 4),
- 'jumping': ('#FF1717', 5),
- 'escape': ('#FF1717', 4),
- 'freezing': ('#FF75EF', 4),
- 'dwelling': ('#FF75EF', 4),
- 'grooming': ('#28AE61', 4),
- 'real_rearing': ('#C409FF', 3),
- 'stretching': ('#C409FF', 3),
- 'nest': ('gray', 2),
- 'others': ('white', 2),
- 'rearing': ('white', 1),
- 'sniffing': ('#0C8140', 1),
- 'looming': ('white', 1),
- 'reaction': ('white', 1),
- 'climbing': ('gray', 1),
- 'in_proximity': ('white', 1)
- }
- fig, axes = plt.subplots(1, 2, figsize=(10, 2), dpi=300)
- plt.subplots_adjust(wspace=0)
- for i, group in enumerate(lst_bl_etho_dict.keys()):
- ax = axes[i]
- ax.set_yticks(range(len(lst_bl_etho_dict[group].keys())))
- for idx, t_id in enumerate(lst_bl_etho_dict[group].keys()):
- behaviors = lst_bl_etho_dict[group][t_id]
- for b, frame_ranges in behaviors.items():
- for frame_range in frame_ranges:
- if np.isnan(frame_ranges).all():
- continue
- start_frame, end_frame = frame_range
- color, zorder = lst_behavior_properties.get(b, ('black', 0))
- color = adjust_saturation(color, saturation)
- rect = plt.Rectangle((start_frame, idx - 0.4), end_frame - start_frame, 0.8, facecolor=color, edgecolor='none', zorder=zorder)
- ax.add_patch(rect)
- if idx % 2 == 0:
- ax.axhline(idx - 0.5, color='black', linestyle='-', linewidth=1, zorder=9)
- ytick_color = 'red'
- else:
- ax.axhline(idx - 0.5, color='black', linestyle='-', linewidth=0.1, zorder=9)
- ytick_color = 'green'
- ax.get_yticklabels()[idx].set_color(ytick_color)
- if i == 0:
- ax.set_yticklabels(['D', 'S', 'D', 'S', 'D', 'S'])
- ax.set_ylabel('Looming')
- else:
- ax.set_yticks([])
- ax.set_xticks([])
- lst_blank_width = (lst_framerange[1] - lst_framerange[0]) / 20
- per_rect = plt.Rectangle((lst_framerange[0] - lst_blank_width, -0.5), lst_blank_width, 6,
- facecolor='white', edgecolor='none', zorder=8)
- ax.add_patch(per_rect)
- post_rect = plt.Rectangle((lst_framerange[1], -0.5), lst_blank_width, 6,
- facecolor='white', edgecolor='none', zorder=8)
- ax.add_patch(post_rect)
- ax.set_xlim(lst_framerange[0] - lst_blank_width, lst_framerange[1] + lst_blank_width)
- ax.set_ylim(-0.5, 5.5)
- ax.invert_yaxis()
- ax.set_title(f'{group}')
- if i == 1:
- line_x1 = lst_framerange[1] - 30*30
- line_x2 = lst_framerange[1]
- ax.text((line_x1 + line_x2) / 2, -0.8, '30 sec', fontsize=8, color='black', ha='center')
- lst_line = plt.Line2D([line_x1, line_x2], [-0.65, -0.65], color='black', linewidth=1)
- lst_line.set_clip_on(False)
- ax.add_artist(lst_line)
- lst_legend_ = {
- 'freezing': ('#FF75EF', 4),
- 'tail rattling': ('#FF09FF', 5),
- 'stretching / rearing': ('#C409FF', 3),
- 'huddling': ('#6D57F3', 1),
- 'approaching P (partner)': ('#143FCA', 4),
- 'following P': ('#137CAB', 4),
- 'grooming P': ('#0990FF', 4),
- 'sniffing P': ('#09FFFF', 4),
- 'grooming': ('#28AE61', 4),
- 'sniffing': ('#0C8140', 1),
- 'in nest': ('gray', 2),
- 'other behaviors': ('white', 1)}
- lst_legend = []
- for label, (color, zorder) in lst_legend_.items():
- if label == 'other behaviors':
- rect = plt.Rectangle((0, 0), 1, 1, facecolor=adjust_saturation(color, saturation),
- edgecolor='black', linewidth=0.5, label=label)
- else:
- rect = plt.Rectangle((0, 0), 1, 1, facecolor=adjust_saturation(color, saturation), label=label)
- lst_legend.append(rect)
- axes[1].legend(handles=lst_legend, loc='upper center', bbox_to_anchor=(0, -0.2), fontsize=7, ncol=7)
- plt.savefig('fig/fig_1/lst_baseline_behavior_ethogram.eps', format="eps", dpi=300, bbox_inches="tight")
- plt.show()
- # %% [markdown]
- # ### FigS1D. Co-trigger percentage
- # %%
- import pandas as pd
- lst_decision_file = os.path.abspath(os.path.join('.', 'lst', 'looming_behavior_sb_1looming_new.csv'))
- trial_data = pd.read_csv(lst_decision_file, header=None, nrows=340)
- # Groups where M1 is MS and M2 is MD
- md_is_m2_groups = ['1', '3', '7', '10', '11', '12', '13', '14', '15', '16', '17', '20']
- md_trigger_count = 0 # only MD triggered
- ms_trigger_count = 0 # only MS triggered
- co_trigger_count = 0 # both triggered
- no_trigger_count = 0 # neither triggered
- for i in range(len(trial_data[0])):
- t_id = trial_data[0][i]
- # Only process paired sessions (no 'M1' or 'M2' in t_id)
- if 'M1' in t_id or 'M2' in t_id:
- continue
- m1_behavior = trial_data[4][i]
- m2_behavior = trial_data[5][i]
- m1_trigger_value = trial_data[6][i]
- m2_trigger_value = trial_data[7][i]
- group_num = t_id.split('B')[1].split('D')[0]
- if group_num in md_is_m2_groups:
- md_trigger_value = m2_trigger_value # M2 is MD
- ms_trigger_value = m1_trigger_value # M1 is MS
- md_behavior = m2_behavior
- ms_behavior = m1_behavior
- else:
- md_trigger_value = m1_trigger_value # M1 is MD
- ms_trigger_value = m2_trigger_value # M2 is MS
- md_behavior = m1_behavior
- ms_behavior = m2_behavior
- md_triggered = (md_trigger_value != 0) and (md_behavior != 'N')
- ms_triggered = (ms_trigger_value != 0) and (ms_behavior != 'N')
- if md_triggered and ms_triggered:
- co_trigger_count += 1
- elif md_triggered:
- md_trigger_count += 1
- elif ms_triggered:
- ms_trigger_count += 1
- else:
- no_trigger_count += 1
- total_trials = md_trigger_count + ms_trigger_count + co_trigger_count + no_trigger_count
- print(f"Paired session trial counts:")
- print(f" md-only trigger: {md_trigger_count}")
- print(f" ms-only trigger: {ms_trigger_count}")
- print(f" co-trigger (both): {co_trigger_count}")
- print(f" no trigger: {no_trigger_count}")
- print(f" total: {total_trials}")
- print(f" md-trigger total (md-only + co): {md_trigger_count + co_trigger_count}")
- print(f" ms-trigger total (ms-only + co): {ms_trigger_count + co_trigger_count}")
- df_co_trigger = pd.DataFrame({
- 'category': ['md-only', 'ms-only', 'co-trigger'],
- 'md_count': [md_trigger_count, 0, co_trigger_count],
- 'ms_count': [0, ms_trigger_count, co_trigger_count]
- })
- df_co_trigger.to_csv('data/FigS1D_co_trigger.csv', index=False)
- # %%
- import numpy as np
- import matplotlib.pyplot as plt
- df = pd.read_csv('data/FigS1D_co_trigger.csv')
- dom_vals = [df.loc[0, 'md_count'], df.loc[2, 'md_count']]
- sub_vals = [df.loc[1, 'ms_count'], df.loc[2, 'ms_count']]
- labels = ['single-trigger', 'co-trigger']
- x = np.array([0, 1]) # Dom, Sub
- width = 0.6
- fig, ax = plt.subplots()
- bottom_dom = 0
- bottom_sub = 0
- ax.bar(0, dom_vals[0], width, bottom=bottom_dom, color='blue')
- ax.bar(0, dom_vals[1], width, bottom=dom_vals[0], color='orange')
- ax.bar(1, sub_vals[0], width, bottom=bottom_sub, color='blue')
- ax.bar(1, sub_vals[1], width, bottom=sub_vals[0], color='orange')
- ax.set_xticks(x)
- ax.set_xticklabels(['Dom', 'Sub'])
- ax.set_ylim(0, 80)
- ax.set_ylabel('Trials')
- plt.tight_layout()
- plt.show()
- # %% [markdown]
- # ### FigS1E. 3 type of grooming time (%)
- # %%
- from analyze_data_utils import filter_dict_data, filter_in_range, calculate_duration_time, calculate_total, calculate_duration_pct, dict_to_dataframe
- self_grooming_data = filter_dict_data(lst_bhvr_labels_dict, 'grooming')
- self_grooming_data = filter_in_range(self_grooming_data, [750,1950], method='replace')
- self_grooming_durations = calculate_duration_time(self_grooming_data)
- total_self_grooming_duration = calculate_total(self_grooming_durations)
- self_grooming_pct = calculate_duration_pct(total_self_grooming_duration, time_range=[20, 60])
- avg_self_grooming_duration_df = dict_to_dataframe(self_grooming_pct, value_name='self_grooming_pct', nan2zero=True)
- # avg_self_grooming_duration_df.to_csv('data/fig1I_suppl_avg_self_grooming_duration.csv', index=False)
- self_grooming_pct_means = avg_self_grooming_duration_df.mean(numeric_only=True)
- grooming_partner_data = filter_dict_data(lst_bhvr_labels_dict, 'groom_partner')
- grooming_partner_data = filter_in_range(grooming_partner_data, [750,1950], method='replace')
- grooming_partner_durations = calculate_duration_time(grooming_partner_data)
- total_grooming_partner_duration = calculate_total(grooming_partner_durations)
- grooming_partner_pct = calculate_duration_pct(total_grooming_partner_duration, time_range=[20, 60])
- avg_grooming_partner_duration_df = dict_to_dataframe(grooming_partner_pct, value_name='grooming_given_pct', nan2zero=True)
- grooming_given_pct_means = avg_grooming_partner_duration_df.mean(numeric_only=True)
- # Calculate grooming received percentages by swapping the given percentages
- grooming_received_pct_means = pd.Series({
- 'grooming_given_pct_sD': grooming_given_pct_means['grooming_given_pct_sS'],
- 'grooming_given_pct_sS': grooming_given_pct_means['grooming_given_pct_sD'],
- 'grooming_given_pct_pD': grooming_given_pct_means['grooming_given_pct_pS'],
- 'grooming_given_pct_pS': grooming_given_pct_means['grooming_given_pct_pD']
- })
- grooming_summary = pd.DataFrame({
- 'self-grooming': [
- self_grooming_pct_means['self_grooming_pct_sD'],
- self_grooming_pct_means['self_grooming_pct_pD'],
- self_grooming_pct_means['self_grooming_pct_sS'],
- self_grooming_pct_means['self_grooming_pct_pS']
- ],
- 'grooming-received': [
- grooming_received_pct_means['grooming_given_pct_sD'],
- grooming_received_pct_means['grooming_given_pct_pD'],
- grooming_received_pct_means['grooming_given_pct_sS'],
- grooming_received_pct_means['grooming_given_pct_pS']
- ],
- 'grooming-given': [
- grooming_given_pct_means['grooming_given_pct_sD'],
- grooming_given_pct_means['grooming_given_pct_pD'],
- grooming_given_pct_means['grooming_given_pct_sS'],
- grooming_given_pct_means['grooming_given_pct_pS']
- ]
- }, index=['sD', 'pD', 'sS', 'pS'])
- grooming_summary.to_csv('data/FigS1E_grooming_3type.csv')
- # %%
- import pandas as pd
- import numpy as np
- import matplotlib.pyplot as plt
- df = pd.read_csv('data/FigS1E_grooming_3type.csv')
- df = df.set_index('Unnamed: 0').loc[['sD','pD','sS','pS']]
- x_labels = ['sD','pD','sS','pS']
- x = np.arange(len(x_labels))
- width = 0.6
- colors = {
- 'self-grooming': 'blue',
- 'grooming-received': 'green',
- 'grooming-given': 'orange'
- }
- fig, ax = plt.subplots()
- bottom = np.zeros(len(x_labels))
- for col in ['self-grooming', 'grooming-received', 'grooming-given']:
- ax.bar(x, df[col].values, width, bottom=bottom, color=colors[col])
- bottom += df[col].values
- ax.set_xticks(x)
- ax.set_xticklabels(x_labels)
- ax.set_ylim(0, 15)
- ax.set_ylabel('Percentage')
- plt.tight_layout()
- plt.show()
- # %% [markdown]
- # ## Fig2
- # %% [markdown]
- # ### Fig2B&D. location and velocity histogram
- # %%
- import numpy as np
- import matplotlib.pyplot as plt
- from matplotlib.colors import ListedColormap, BoundaryNorm
- import seaborn as sns
- from analyze_data_utils import ret_extract_target_data
- import pickle
- import pandas as pd
- def prepare_velocity_data(categories, groups, ranks, data_source, min_speed=0, max_speed=300):
- all_x = []
- all_y = []
- all_vx = []
- all_vy = []
- all_ids = []
- time = 'withrat'
- x_offset = 12
- for categorie in categories:
- for group in groups:
- for rank in ranks:
- with open ("G:\Li_lab\ppt\S_paper\Paper_v3\check_ret_zone\zone_correct.pkl", 'rb') as f:
- zone_correct = pickle.load(f)
- m_id = group+rank
- dx = (zone_correct['w_correct'][m_id] - 0.5) / 24
- rat_x = zone_correct['x_correct'][m_id] - x_offset
- x_min, x_max = rat_x, rat_x+(24+x_offset)*dx
- y_min, y_max = -10.5, 9.5
- interested_frame = [0, 8900]
- x_center_bl = ret_extract_target_data(ret_raw_data, categories=[categorie],
- groups=[group], ranks=[rank],
- target_key=['X center'], times=[time])
- y_center_bl = ret_extract_target_data(ret_raw_data, categories=[categorie],
- groups=[group], ranks=[rank],
- target_key=['Y center'], times=[time])
- # velocity_al = ret_extract_target_data(ret_raw_data, categories=[categorie],
- # groups=[group], ranks=[rank],
- # target_key=['Velocity'], times=['withrat'])
- # data_to_save = {'x_center_al': x_center_bl, 'y_center_al': y_center_bl, 'velocity_al': velocity_al}
- # with open(r'G:\Li_lab\ppt\S_paper\Paper_v3\check_ret_zone\test_center_data.pkl', 'wb') as f:
- # pickle.dump(data_to_save, f)
- x = x_center_bl[categorie][rank][time][group].iloc[interested_frame[0]:interested_frame[1]]
- y = y_center_bl[categorie][rank][time][group].iloc[interested_frame[0]:interested_frame[1]]
- if m_id == 'D4MD':
- x = x[~x.index.isin([7181, 7182])]
- y = y[~y.index.isin([7181, 7182])]
- x = pd.to_numeric(x, errors='coerce')
- y = pd.to_numeric(y, errors='coerce')
- x = np.array(x, dtype=float)
- y = np.array(y, dtype=float)
- vx = np.diff(x) * 30
- vy = np.diff(y) * 30
- speed = np.sqrt(vx**2 + vy**2)
- mask = (speed >= min_speed) & (speed <= max_speed) & \
- (x[:-1] >= x_min) & (x[:-1] < x_max) & \
- (y[:-1] >= y_min) & (y[:-1] < y_max)
- all_x.extend(x[:-1][mask])
- all_y.extend(y[:-1][mask])
- all_vx.extend(vx[mask])
- all_vy.extend(vy[mask])
- all_ids.extend([group+rank] * np.sum(mask))
- # Convert to NumPy arrays
- all_x = np.array(all_x)
- all_y = np.array(all_y)
- all_vx = np.array(all_vx)
- all_vy = np.array(all_vy)
- all_ids = np.array(all_ids)
- return all_x, all_y, all_vx, all_vy, all_ids
- groups = ['D1', 'D2', 'D3', 'D4', 'D5', 'D6', 'D7', 'D8', 'D9', 'D10', 'D11']
- ranks = ['MD', 'MS']
- xbins = 12
- ybins = 10
- x, y, vx, vy, point_ids = prepare_velocity_data(['single'], groups, ranks, ret_raw_data)
- count, xedges, yedges = np.histogram2d(x, y, bins=[xbins, ybins])
- mouse_count_grid = np.zeros_like(count, dtype=int)
- x_bin_indices = np.digitize(x, xedges) - 1
- y_bin_indices = np.digitize(y, yedges) - 1
- cell_mice_dict = {}
- for xi, yi, mouse_id in zip(x_bin_indices, y_bin_indices, point_ids):
- if 0 <= xi < xbins and 0 <= yi < ybins:
- key = (xi, yi)
- if key not in cell_mice_dict:
- cell_mice_dict[key] = set()
- cell_mice_dict[key].add(mouse_id)
- for (xi, yi), mice_set in cell_mice_dict.items():
- mouse_count_grid[xi, yi] = len(mice_set)
- mask = (mouse_count_grid >= 2) & (count >= 5)
- grid_vx, _, _ = np.histogram2d(x, y, bins=[xedges, yedges], weights=vx)
- grid_vy, _, _ = np.histogram2d(x, y, bins=[xedges, yedges], weights=vy)
- grid_vx /= count
- grid_vy /= count
- grid_vx[~mask] = 1e-5
- grid_vy[~mask] = 1e-5
- speed_grid = np.sqrt(grid_vx**2 + grid_vy**2)
- speeds = np.sqrt(vx**2 + vy**2)
- speed_sum, _, _ = np.histogram2d(x, y, bins=[xedges, yedges], weights=speeds)
- speed_avg = speed_sum / count
- speed_avg[~mask] = 1e-5
- heatmap, xedges, yedges = np.histogram2d(x, y, bins=[xedges, yedges])
- hist_frequency = heatmap / np.sum(heatmap) * 100
- hist_frequency[~mask] = 1e-5
- df_grid = pd.DataFrame({
- 'x_index': np.tile(np.arange(xbins), ybins),
- 'y_index': np.repeat(np.arange(ybins), xbins),
- 'time_pct': hist_frequency.T.flatten(),
- 'speed_avg': speed_avg.T.flatten(),
- 'vx': grid_vx.T.flatten(),
- 'vy': grid_vy.T.flatten()
- })
- df_grid.to_csv('data/Fig2B&D_location_speed_heatmap.csv', index=False)
- # %%
- import numpy as np
- import matplotlib.pyplot as plt
- import pandas as pd
- # Load data
- df = pd.read_csv('data/Fig2B&D_location_speed_heatmap.csv')
- xbins = 12
- ybins = 10
- # Reshape to matrices (shape: (ybins, xbins))
- hist_frequency = df.pivot(index='y_index', columns='x_index', values='time_pct').values
- speed_avg = df.pivot(index='y_index', columns='x_index', values='speed_avg').values
- grid_vx = df.pivot(index='y_index', columns='x_index', values='vx').values
- grid_vy = df.pivot(index='y_index', columns='x_index', values='vy').values
- # Figure 1: Coordinate heatmap
- fig, ax = plt.subplots(figsize=(5,5), dpi=200)
- im = ax.imshow(
- hist_frequency, # no transpose (shape: 10x12)
- cmap='viridis',
- origin='lower',
- extent=(0, xbins, 0, ybins),
- vmin=0, vmax=12
- )
- fig.colorbar(im, label='Time (%)')
- ax.set_xlim(0, xbins)
- ax.set_ylim(0, ybins)
- ax.set_xticks([])
- ax.set_yticks([])
- ax.invert_yaxis()
- plt.tight_layout()
- # plt.savefig(...)
- plt.show()
- # Figure 2: Speed heatmap with arrows
- fig, ax = plt.subplots(figsize=(5,5), dpi=200)
- # Speed heatmap
- im2 = ax.imshow(
- speed_avg, # no transpose
- cmap='viridis',
- origin='lower',
- extent=(0, xbins, 0, ybins),
- vmin=0, vmax=20
- )
- plt.colorbar(im2, label='Speed (cm/s)')
- # Arrows
- X, Y = np.meshgrid(np.arange(xbins) + 0.5, np.arange(ybins) + 0.5)
- scale = 0.3
- head_scale = 0.1
- for i in range(X.shape[0]): # y index (0 ~ ybins-1)
- for j in range(X.shape[1]): # x index (0 ~ xbins-1)
- vx_ = grid_vx[i, j]
- vy_ = grid_vy[i, j]
- mag = np.sqrt(vx_**2 + vy_**2)
- dx = vx_ * scale
- dy = vy_ * scale
- x0 = X[i, j]
- y0 = Y[i, j]
- x1 = x0 + dx
- y1 = y0 + dy
- ax.annotate('', xy=(x1, y1), xytext=(x0, y0),
- arrowprops=dict(
- arrowstyle='->, head_width={:.2f}, head_length={:.2f}'.format(
- mag * head_scale, mag * head_scale * 1.5
- ),
- color='white',
- linewidth=1,
- mutation_scale=5,
- shrinkA=0, shrinkB=0
- ))
- ax.invert_yaxis()
- ax.set_xlim(0, xbins)
- ax.set_ylim(0, ybins)
- ax.set_xticks([])
- ax.set_yticks([])
- plt.tight_layout()
- # plt.savefig(...)
- plt.show()
- # %% [markdown]
- # ### Fig2C&E & 2C suppl: Time in zone, Speed in zone (ret, withrat 0-8900)
- # %%
- import numpy as np
- import pandas as pd
- from analyze_data_utils import ret_extract_target_data, get_ret_location_value
- # Define parameters
- fps = 30
- max_frames = 8900 # withrat 0-8900
- # Arena parameters (consistent with other ret analyses)
- dx = (19 + 5) / 24 # cm per unit (= 1.0)
- rat_x = -5 # rat side X coordinate start (cm)
- arena_size = 24 * dx # arena total length 24 cm
- # Zone definitions: near:middle:far = 9:9:6 (cm)
- near_zone_width = 9 * dx
- middle_zone_width = 9 * dx
- far_zone_width = 6 * dx
- near_zone_range = (rat_x, rat_x + near_zone_width) # (-5, 4)
- middle_zone_range = (rat_x + near_zone_width, rat_x + near_zone_width + middle_zone_width) # ( 4, 13)
- far_zone_range = (rat_x + near_zone_width + middle_zone_width, rat_x + arena_size) # (13, 19)
- zone_names = {0: 'near', 1: 'middle', 2: 'far'}
- ret_groups = [f'D{i}' for i in range(1, 12)] # D1 ~ D11
- # Extract X center and Y center (withrat period)
- x_center_ret = ret_extract_target_data(ret_raw_data, target_key=['X center'], times=['withrat'])
- y_center_ret = ret_extract_target_data(ret_raw_data, target_key=['Y center'], times=['withrat'])
- # ret_zones_time_pcts[condition][rank][zone] = [value for D1..D11]
- ret_zones_time_pcts = {}
- ret_zones_speed_vals = {}
- for cat in ['single', 'pair']:
- ret_zones_time_pcts[cat] = {'MD': {0: [], 1: [], 2: []}, 'MS': {0: [], 1: [], 2: []}}
- ret_zones_speed_vals[cat] = {'MD': {0: [], 1: [], 2: []}, 'MS': {0: [], 1: [], 2: []}}
- for cat in ['single', 'pair']:
- for rank in ['MD', 'MS']:
- for d in ret_groups:
- try:
- x_series = x_center_ret[cat][rank]['withrat'][d]
- y_series = y_center_ret[cat][rank]['withrat'][d]
- except (KeyError, TypeError):
- for zone in [0, 1, 2]:
- ret_zones_time_pcts[cat][rank][zone].append(np.nan)
- ret_zones_speed_vals[cat][rank][zone].append(np.nan)
- continue
- # Coordinate time correction: align behavior frames with coordinate frames using rat_in_frames and offset_frames
- # Determine correct session_id based on condition and rank
- if cat == 'pair':
- sess_id = d + 'M1&M2'
- else: # cat == 'single'
- day_num = d.split('D')[1].split('M')[0] # Extract number, e.g., 'D4' -> '4'
- if day_num == '4': # D4 special case
- m_id = '1' if rank == 'MS' else '2'
- else: # all other sessions
- m_id = '1' if rank == 'MD' else '2'
- sess_id = d + 'M' + m_id
- # Calculate offset
- corr_rat_in = rat_in_frames[sess_id] - offset_frames[sess_id]
- x_raw = pd.to_numeric(x_series, errors='coerce')
- y_raw = pd.to_numeric(y_series, errors='coerce')
- # Handle negative offset: behavior data starts earlier than coordinate data
- if corr_rat_in >= 0:
- # Normal case: take max_frames frames starting from corr_rat_in
- x = x_raw.iloc[corr_rat_in:corr_rat_in + max_frames].reset_index(drop=True)
- y = y_raw.iloc[corr_rat_in:corr_rat_in + max_frames].reset_index(drop=True)
- else:
- # Negative offset: pad NaN at the beginning, then take data from 0
- n_pad = -corr_rat_in # Number of frames to pad
- n_data = max_frames - n_pad # Number of frames to take from coordinate data
- # Create padded NaN series
- pad_series = pd.Series([np.nan] * n_pad)
- # Extract coordinate data (starting from 0, take n_data frames)
- x_data = x_raw.iloc[:n_data]
- y_data = y_raw.iloc[:n_data]
- # Concatenate
- x = pd.concat([pad_series, x_data], ignore_index=True)
- y = pd.concat([pad_series, y_data], ignore_index=True)
- # Compute zone labels
- location_value = get_ret_location_value(x)
- # Compute frame-by-frame speed (cm/s) from X/Y displacement
- dx_pos = x.diff().fillna(0)
- dy_pos = y.diff().fillna(0)
- speed = np.sqrt(dx_pos**2 + dy_pos**2) * fps
- for zone in [0, 1, 2]:
- zone_mask = (location_value == zone)
- zone_count = zone_mask.sum()
- zone_pct = zone_count / max_frames * 100
- ret_zones_time_pcts[cat][rank][zone].append(zone_pct)
- zone_speed_mean = speed[zone_mask].mean() if zone_mask.any() else np.nan
- ret_zones_speed_vals[cat][rank][zone].append(zone_speed_mean)
- # %%
- import pandas as pd
- import numpy as np
- zone_order = ['near', 'middle', 'far']
- zone_id_map = {'near': 0, 'middle': 1, 'far': 2}
- ret_groups = [f'D{i}' for i in range(1, 12)]
- # -- Time percentage DataFrame -----------------------------------------
- df_ret_zones_pct = {'single': None, 'pair': None}
- df_ret_zones_unit = {'single': None, 'pair': None}
- for cat in ['single', 'pair']:
- data_pct = []
- for rank in ['MD', 'MS']:
- for i, d in enumerate(ret_groups):
- row = {'D_id': d, 'rank': rank}
- for zname in zone_order:
- vals = ret_zones_time_pcts[cat][rank][zone_id_map[zname]]
- row[f'{zname}_pct'] = vals[i] if i < len(vals) else np.nan
- data_pct.append(row)
- df_ret_zones_pct[cat] = pd.DataFrame(data_pct)
- # Normalize: % -> s/(min·cm²)
- # withrat duration (min)
- withrat_time_min = max_frames / fps / 60
- # Zone areas (cm²): arena width 24 cm
- zone_areas = {'near': 9 * 24, 'middle': 9 * 24, 'far': 6 * 24}
- df_ret_zones_unit[cat] = df_ret_zones_pct[cat].copy()
- for zname in zone_order:
- col = f'{zname}_pct'
- # pct / 100 * total_s / time_min / area_cm2 = s/(min·cm²)
- df_ret_zones_unit[cat][col] = (
- df_ret_zones_pct[cat][col] / 100 * 60 / zone_areas[zname]
- )
- # Adjust column order
- # Info columns
- info_cols = ['D_id', 'rank']
- # Data columns
- data_cols = [f'{zone}_pct' for zone in zone_order]
- for cat in ['single', 'pair']:
- df_ret_zones_pct[cat] = df_ret_zones_pct[cat][info_cols + data_cols]
- df_ret_zones_unit[cat] = df_ret_zones_unit[cat][info_cols + data_cols]
- # Save DataFrame (only save single condition)
- df_ret_zones_pct['single'].to_csv('data/FigS2C_zones_time_pcts.csv', index=False)
- df_ret_zones_unit['single'].to_csv('data/Fig2C_zones_time_unit.csv', index=False)
- # -- Speed DataFrame -----------------------------------------------
- df_ret_zones_speed = {'single': None, 'pair': None}
- for cat in ['single', 'pair']:
- data_spd = []
- for rank in ['MD', 'MS']:
- for i, d in enumerate(ret_groups):
- row = {'D_id': d, 'rank': rank}
- for zname in zone_order:
- vals = ret_zones_speed_vals[cat][rank][zone_id_map[zname]]
- row[f'{zname}_speed'] = vals[i] if i < len(vals) else np.nan
- data_spd.append(row)
- df_ret_zones_speed[cat] = pd.DataFrame(data_spd)
- # Adjust column order
- # Info columns
- info_cols = ['D_id', 'rank']
- # Data columns
- data_cols = [f'{zone}_speed' for zone in zone_order]
- for cat in ['single', 'pair']:
- df_ret_zones_speed[cat] = df_ret_zones_speed[cat][info_cols + data_cols]
- # Save DataFrame (only save single condition)
- df_ret_zones_speed['single'].to_csv('data/Fig2E_zones_speed.csv', index=False)
- # %%
- df = pd.read_csv('data/FigS2C_zones_time_pcts.csv')
- print(df.columns.tolist())
- exclude_cols = ['B_id', 'rank', 'mouse_id']
- plot_cols = [col for col in df.columns if col not in exclude_cols]
- means = df[plot_cols].mean()
- plt.figure()
- bar_colors = ['orange', 'blue', 'green']
- bars = plt.bar(means.index, means.values, color=bar_colors[:len(means)])
- df_melted = df[plot_cols]
- sns.stripplot(
- data=df_melted,
- jitter=True,
- color='gray',
- size=5
- )
- plt.ylabel('Time in zone (%)')
- plt.xticks(rotation=45, ha='right')
- plt.gca().invert_xaxis()
- plt.ylim(0, 100)
- plt.show()
- # %%
- df = pd.read_csv('data/Fig2C_zones_time_unit.csv')
- print(df.columns.tolist())
- exclude_cols = ['B_id', 'rank', 'mouse_id']
- plot_cols = [col for col in df.columns if col not in exclude_cols]
- means = df[plot_cols].mean()
- plt.figure()
- bar_colors = ['orange', 'blue', 'green']
- bars = plt.bar(means.index, means.values, color=bar_colors[:len(means)])
- df_melted = df[plot_cols]
- sns.stripplot(
- data=df_melted,
- jitter=True,
- color='gray',
- size=5
- )
- plt.ylabel('Time in zone (s/(min﹡cm²))')
- plt.xticks(rotation=45, ha='right')
- plt.gca().invert_xaxis()
- plt.ylim(0, 0.4)
- plt.show()
- # %%
- df = pd.read_csv('data/Fig2E_zones_speed.csv')
- print(df.columns.tolist())
- exclude_cols = ['B_id', 'rank', 'mouse_id']
- plot_cols = [col for col in df.columns if col not in exclude_cols]
- means = df[plot_cols].mean()
- plt.figure()
- bar_colors = ['orange', 'blue', 'green']
- bars = plt.bar(means.index, means.values, color=bar_colors[:len(means)])
- df_melted = df[plot_cols]
- sns.stripplot(
- data=df_melted,
- jitter=True,
- color='gray',
- size=5
- )
- plt.ylabel('Speed in zone (cm/s)')
- plt.xticks(rotation=45, ha='right')
- plt.gca().invert_xaxis()
- plt.ylim(0, 9)
- plt.show()
- # %% [markdown]
- # ### Fig2F. ΔTime in zone (%)
- # %%
- import pandas as pd
- import numpy as np
- # Define zone order
- zone_order = ['near', 'middle', 'far']
- ret_groups = [f'D{i}' for i in range(1, 12)]
- delta_data = []
- for d_id in ret_groups:
- row = {'D_id': d_id}
- for rank in ['MD', 'MS']:
- single_row = df_ret_zones_pct['single'][
- (df_ret_zones_pct['single']['D_id'] == d_id) &
- (df_ret_zones_pct['single']['rank'] == rank)
- ]
- pair_row = df_ret_zones_pct['pair'][
- (df_ret_zones_pct['pair']['D_id'] == d_id) &
- (df_ret_zones_pct['pair']['rank'] == rank)
- ]
- for zone in zone_order:
- col = f'{zone}_pct'
- if len(single_row) > 0 and len(pair_row) > 0:
- row[f'{zone}_{rank}'] = pair_row[col].values[0] - single_row[col].values[0]
- else:
- row[f'{zone}_{rank}'] = np.nan
- delta_data.append(row)
- # Create DataFrame and adjust column order
- df_delta_time = pd.DataFrame(delta_data)
- df_delta_time = df_delta_time[['D_id', 'near_MD', 'near_MS', 'middle_MD', 'middle_MS', 'far_MD', 'far_MS']]
- df_delta_time.columns = [col.replace('_MD', '_D').replace('_MS', '_S') for col in df_delta_time.columns]
- # Save result
- df_delta_time.to_csv('data/Fig2F_delta_zones_time.csv', index=False)
- # %%
- df = pd.read_csv(f'data/Fig2F_delta_zones_time.csv')
- print(df.columns.tolist())
- groups = ['far', 'middle', 'near']
- fig, ax = plt.subplots()
- x_pos = {}
- xticks = []
- xticklabels = []
- x = 0
- colors = {'D': 'orange', 'S': 'blue'}
- # 1. x
- for g in groups:
- for cond in ['D', 'S']:
- if cond == 'D':
- x_pos[f'{g}_{cond}_mean'] = x
- x_pos[f'{g}_{cond}_pts'] = x + 0.8
- xticks += [x, x + 0.8]
- xticklabels += [f'{g}_{cond}_mean', f'{g}_{cond}_pts']
- else:
- x_pos[f'{g}_{cond}_pts'] = x
- x_pos[f'{g}_{cond}_mean'] = x + 0.8
- xticks += [x, x + 0.8]
- xticklabels += [f'{g}_{cond}_pts', f'{g}_{cond}_mean']
- x += 2.2
- # 2. mean ± SD
- for g in groups:
- for cond in ['D', 'S']:
- col = f'{g}_{cond}'
- xpos = x_pos[f'{g}_{cond}_mean']
- ax.errorbar(
- xpos,
- df[col].mean(),
- yerr=df[col].std(),
- fmt='o',
- color=colors[cond],
- capsize=4,
- markersize=9, # 固定大小
- elinewidth=1.5,
- capthick=1.5,
- zorder=3
- )
- # 3. raw points + pairing
- for i in range(len(df)):
- for g in groups:
- D_col = f'{g}_D'
- S_col = f'{g}_S'
- x_d = x_pos[f'{g}_D_pts']
- y_d = df.loc[i, D_col]
- x_s = x_pos[f'{g}_S_pts']
- y_s = df.loc[i, S_col]
- ax.scatter(x_d, y_d, color=colors['D'])
- ax.scatter(x_s, y_s, color=colors['S'])
- ax.plot([x_d, x_s], [y_d, y_s],
- color='gray', linewidth=1)
- # 4. axis
- ax.set_xticks(xticks)
- ax.set_xticklabels(xticklabels, rotation=45, ha='right')
- ax.set_ylabel('Δ Time in zone (%)')
- ax.set_ylim([-60, 40])
- plt.tight_layout()
- plt.show()
- # %% [markdown]
- # ### Fig2G. Δ Speed in zone (cm/s)
- # %%
- import pandas as pd
- import numpy as np
- # Define zone order
- zone_order = ['near', 'middle', 'far']
- ret_groups = [f'D{i}' for i in range(1, 12)]
- delta_data = []
- for d_id in ret_groups:
- row = {'D_id': d_id}
- for rank in ['MD', 'MS']:
- single_row = df_ret_zones_speed['single'][
- (df_ret_zones_speed['single']['D_id'] == d_id) &
- (df_ret_zones_speed['single']['rank'] == rank)
- ]
- pair_row = df_ret_zones_speed['pair'][
- (df_ret_zones_speed['pair']['D_id'] == d_id) &
- (df_ret_zones_speed['pair']['rank'] == rank)
- ]
- for zone in zone_order:
- col = f'{zone}_speed'
- if len(single_row) > 0 and len(pair_row) > 0:
- row[f'{zone}_{rank}'] = pair_row[col].values[0] - single_row[col].values[0]
- else:
- row[f'{zone}_{rank}'] = np.nan
- delta_data.append(row)
- # Create DataFrame and adjust column order
- df_delta_speed = pd.DataFrame(delta_data)
- df_delta_speed = df_delta_speed[['D_id', 'near_MD', 'near_MS', 'middle_MD', 'middle_MS', 'far_MD', 'far_MS']]
- df_delta_speed.columns = [col.replace('_MD', '_D').replace('_MS', '_S') for col in df_delta_speed.columns]
- df_delta_speed.to_csv('data/Fig2G_delta_zones_speed.csv', index=False)
- # %%
- df = pd.read_csv(f'data/Fig2G_delta_zones_speed.csv')
- print(df.columns.tolist())
- groups = ['far', 'middle', 'near']
- fig, ax = plt.subplots()
- x_pos = {}
- xticks = []
- xticklabels = []
- x = 0
- colors = {'D': 'orange', 'S': 'blue'}
- # 1. x
- for g in groups:
- for cond in ['D', 'S']:
- if cond == 'D':
- x_pos[f'{g}_{cond}_mean'] = x
- x_pos[f'{g}_{cond}_pts'] = x + 0.8
- xticks += [x, x + 0.8]
- xticklabels += [f'{g}_{cond}_mean', f'{g}_{cond}_pts']
- else:
- x_pos[f'{g}_{cond}_pts'] = x
- x_pos[f'{g}_{cond}_mean'] = x + 0.8
- xticks += [x, x + 0.8]
- xticklabels += [f'{g}_{cond}_pts', f'{g}_{cond}_mean']
- x += 2.2
- # 2. mean ± SD
- for g in groups:
- for cond in ['D', 'S']:
- col = f'{g}_{cond}'
- xpos = x_pos[f'{g}_{cond}_mean']
- ax.errorbar(
- xpos,
- df[col].mean(),
- yerr=df[col].std(),
- fmt='o',
- color=colors[cond],
- capsize=4,
- markersize=9,
- elinewidth=1.5,
- capthick=1.5,
- zorder=3
- )
- # 3. raw points + pairing
- for i in range(len(df)):
- for g in groups:
- D_col = f'{g}_D'
- S_col = f'{g}_S'
- x_d = x_pos[f'{g}_D_pts']
- y_d = df.loc[i, D_col]
- x_s = x_pos[f'{g}_S_pts']
- y_s = df.loc[i, S_col]
- ax.scatter(x_d, y_d, color=colors['D'])
- ax.scatter(x_s, y_s, color=colors['S'])
- ax.plot([x_d, x_s], [y_d, y_s],
- color='gray', linewidth=1)
- # 4. axis
- ax.set_xticks(xticks)
- ax.set_xticklabels(xticklabels, rotation=45, ha='right')
- ax.set_ylabel('Δ Speed in zone (cm/s)')
- ax.set_ylim([-4, 6])
- plt.tight_layout()
- plt.show()
- # %% [markdown]
- # ### Fig2H. with rat ethogram example
- # %%
- import numpy as np
- import pandas as pd
- import matplotlib.pyplot as plt
- import matplotlib.colors as mcolors
- import colorsys
- from matplotlib.patches import FancyArrow
- from analyze_data_utils import filter_in_range
- def adjust_saturation(color, saturation):
- rgb = mcolors.to_rgb(color)
- h, l, s = colorsys.rgb_to_hls(*rgb)
- new_s = saturation * s
- new_rgb = colorsys.hls_to_rgb(h, l, new_s)
- return new_rgb
- with open('data/ret_etho_dict.pkl', 'rb') as f:
- ret_etho_dict = pickle.load(f)
- # ret_etho_dict = {
- # 'single': {
- # 'D6MD': ret_bhvr_labels_dict['s']['D']['D6MD'],
- # 'D6MS': ret_bhvr_labels_dict['s']['S']['D6MS'],
- # 'D7MD': ret_bhvr_labels_dict['s']['D']['D7MD'],
- # 'D7MS': ret_bhvr_labels_dict['s']['S']['D7MS'],
- # 'D8MD': ret_bhvr_labels_dict['s']['D']['D8MD'],
- # 'D8MS': ret_bhvr_labels_dict['s']['S']['D8MS']
- # },
- # 'pair': {
- # 'D6MD': ret_bhvr_labels_dict['p']['D']['D6MD'],
- # 'D6MS': ret_bhvr_labels_dict['p']['S']['D6MS'],
- # 'D7MD': ret_bhvr_labels_dict['p']['D']['D7MD'],
- # 'D7MS': ret_bhvr_labels_dict['p']['S']['D7MS'],
- # 'D8MD': ret_bhvr_labels_dict['p']['D']['D8MD'],
- # 'D8MS': ret_bhvr_labels_dict['p']['S']['D8MS']
- # }}
- # with open('data/ret_etho_dict.pkl', 'wb') as f:
- # pickle.dump(ret_etho_dict, f)
- ret_framerange = [0, 8900]
- saturation = 0.9
- framerate=30
- stim_frame = int(0.725 * framerate)
- behavior_frames_dict = ret_etho_dict
- behavior_frames_dict = dict(sorted(behavior_frames_dict.items()))
- ret_behavior_frames_dict = filter_in_range(ret_etho_dict, ret_framerange, method='replace')
- # %%
- ret_behavior_properties = {
- 'approach_partner': ('#143FCA', 4),
- 'follow_partner': ('#137CAB', 4),
- 'groom_partner': ('#0990FF', 4),
- 'sniff_partner': ('#09FFFF', 4),
- 'huddling': ('#6D57F3', 4),
- 'approach': ('#FFCA09', 3),
- 'dwelling': ('#FF8409', 3),
- 'withdrawal': ('#FF0990', 3),
- 'stretch_attend': ('#FF7979', 3),
- 'freezing': ('#FF75EF', 3),
- 'tail_rattling': ('#FF09FF', 3),
- "rearing": ("#C409FF", 2),
- 'grooming': ('#28AE61', 2),
- "sniffing": ("#0FD400", 2),
- 'in_proximity': ('white', 1),
- 'rat_in': ('white', 1),
- 'others': ('white', 1)}
- saturation = 0.9
- framerate = 30
- fig, axes = plt.subplots(1, 2, figsize=(10, 2), dpi=300)
- plt.subplots_adjust(wspace=0)
- for i, group in enumerate(ret_etho_dict.keys()):
- ax = axes[i]
- ax.set_yticks(range(len(ret_etho_dict[group].keys())))
- for idx, t_id in enumerate(ret_etho_dict[group].keys()):
- behaviors = ret_etho_dict[group][t_id]
- for b, frame_ranges in behaviors.items():
- for frame_range in frame_ranges:
- if np.isnan(frame_range).all():
- continue
- start_frame, end_frame = frame_range
- color, zorder = ret_behavior_properties.get(b, ('white', 0))
- color = adjust_saturation(color, saturation)
- rect = plt.Rectangle((start_frame, idx - 0.4), end_frame - start_frame, 0.8, facecolor=color, edgecolor='none', zorder=zorder)
- ax.add_patch(rect)
- if idx % 2 == 0:
- ax.axhline(idx - 0.5, color='black', linestyle='-', linewidth=1)
- ytick_color = 'red'
- else:
- ax.axhline(idx - 0.5, color='black', linestyle='--', linewidth=0.1)
- ytick_color = 'green'
- ax.get_yticklabels()[idx].set_color(ytick_color)
- if i == 0:
- ax.set_yticklabels(['D', 'S', 'D', 'S', 'D', 'S'])
- ax.set_ylabel('Rat')
- else:
- ax.set_yticks([])
- ax.set_xticks([])
- ret_mod_width = (ret_framerange[1] - ret_framerange[0]) / 12
- ret_blank_width = (ret_framerange[1] - ret_framerange[0] + ret_mod_width) / 20
- per_rect = plt.Rectangle((ret_framerange[0] - ret_mod_width - ret_blank_width, -0.5), ret_blank_width, 6,
- facecolor='white', edgecolor='white', zorder=8)
- ax.add_patch(per_rect)
- post_rect = plt.Rectangle((ret_framerange[1], -0.5), ret_blank_width, 6,
- facecolor='white', edgecolor='white', zorder=8)
- ax.add_patch(post_rect)
- ax.set_xlim(ret_framerange[0] - ret_blank_width - ret_mod_width, ret_framerange[1] + ret_blank_width)
- ax.set_ylim(-0.5, 5.5)
- ax.invert_yaxis()
- 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')
- ret_arrow.set_clip_on(False)
- ax.add_patch(ret_arrow)
- threat_line = plt.Line2D([0, 8900], [5.75, 5.75], color='red', linewidth=2)
- threat_line.set_clip_on(False)
- ax.add_artist(threat_line)
- if i == 1:
- line_x1 = ret_framerange[1] - 30*60
- line_x2 = ret_framerange[1]
- ax.text((line_x1 + line_x2) / 2, -0.8, '60 sec', fontsize=8, color='black', ha='center')
- ret_line = plt.Line2D([line_x1, line_x2], [-0.65, -0.65], color='black', linewidth=1)
- ret_line.set_clip_on(False)
- ax.add_artist(ret_line)
- ret_legend_ = {
- 'approach': ('#FFCA09', 3),
- 'investigation': ('#FF8409', 3),
- 'withdrawal': ('#FF0990', 3),
- 'stretch-attend': ('#FF7979', 3),
- 'freezing': ('#FF75EF', 3),
- 'tail rattling': ('#FF09FF', 3),
- "rearing": ("#C409FF", 2),
- 'huddling': ('#6D57F3', 1),
- 'approaching P': ('#143FCA', 4),
- 'following P': ('#137CAB', 4),
- 'grooming P': ('#0990FF', 4),
- 'sniffing P': ('#09FFFF', 4),
- 'grooming': ('#28AE61', 2),
- "sniffing": ("#0FD400", 2),
- 'other behaviors': ('white', 1)}
- ret_legend = []
- for label, (color, zorder) in ret_legend_.items():
- if label == 'other behaviors':
- rect = plt.Rectangle((0, 0), 1, 1, facecolor=adjust_saturation(color, saturation),
- edgecolor='black', linewidth=0.5, label=label)
- else:
- rect = plt.Rectangle((0, 0), 1, 1, facecolor=adjust_saturation(color, saturation), label=label)
- ret_legend.append(rect)
- axes[1].legend(handles=ret_legend, loc='upper center', bbox_to_anchor=(0, -0.2), fontsize=7, ncol=7)
- # plt.savefig('fig/fig_2/ret_behavior_ethogram.eps', format="eps", dpi=300, bbox_inches="tight")
- plt.show()
- # %%
- # Assume ret_framerange = (0, 8900) # Infer from your actual situation, from original code
- ret_framerange = (0, 8900) # If not explicitly defined, add this line
- with open('data/ret_bl_etho_dict.pkl', 'rb') as f:
- ret_bl_etho_dict = pickle.load(f)
- # 1. Calculate ret_mod_width and ret_blank_width
- ret_mod_width = (ret_framerange[1] - ret_framerange[0]) / 12 # ≈ 741.67
- ret_blank_width = (ret_framerange[1] - ret_framerange[0] + ret_mod_width) / 20
- # 2. Define baseline window (last ret_mod_width frames)
- baseline_end = 8900 # consistent with ret_framerange[1]
- window_start = baseline_end - ret_mod_width
- window_end = baseline_end
- # 3. Mapping function: baseline frame f -> transition zone new coordinate new_x
- def map_to_transition(f):
- # Map window_end (8900) to ret_framerange[0] (0)
- # Map window_start to ret_framerange[0] - ret_mod_width
- return ret_framerange[0] - ret_mod_width + (f - window_start)
- # Start plotting
- fig, axes = plt.subplots(1, 2, figsize=(10, 2), dpi=300)
- plt.subplots_adjust(wspace=0)
- for i, group in enumerate(ret_etho_dict.keys()):
- ax = axes[i]
- n_animals = len(ret_etho_dict[group])
- ax.set_yticks(range(n_animals))
- # ========== Step 1: Draw baseline transition data (left) ==========
- # Assume ret_bl_etho_dict has the same group keys
- if group in ret_bl_etho_dict:
- bl_group_data = ret_bl_etho_dict[group]
- for idx, t_id in enumerate(bl_group_data.keys()):
- behaviors = bl_group_data[t_id]
- for b, frame_ranges in behaviors.items():
- for frame_range in frame_ranges:
- # Skip invalid values (nan or None)
- try:
- if np.isnan(frame_range).all():
- continue
- except:
- continue
- start_frame, end_frame = frame_range
- # Check overlap with transition window [window_start, window_end]
- if end_frame <= window_start or start_frame >= window_end:
- continue
- # Clip overlap
- clip_start = max(start_frame, window_start)
- clip_end = min(end_frame, window_end)
- # Map to new coordinates
- new_start = map_to_transition(clip_start)
- new_end = map_to_transition(clip_end)
- width = new_end - new_start
- if width <= 0:
- continue
- color, zorder = ret_behavior_properties.get(b, ('white', 0))
- color = adjust_saturation(color, saturation)
- rect = plt.Rectangle((new_start, idx - 0.4), width, 0.8,
- facecolor=color, edgecolor='none', zorder=zorder)
- ax.add_patch(rect)
- # ========== Step 2: Draw original ret_etho_dict data (right) ==========
- for idx, t_id in enumerate(ret_etho_dict[group].keys()):
- behaviors = ret_etho_dict[group][t_id]
- for b, frame_ranges in behaviors.items():
- for frame_range in frame_ranges:
- if np.isnan(frame_range).all():
- continue
- start_frame, end_frame = frame_range
- color, zorder = ret_behavior_properties.get(b, ('white', 0))
- color = adjust_saturation(color, saturation)
- rect = plt.Rectangle((start_frame, idx - 0.4), end_frame - start_frame, 0.8,
- facecolor=color, edgecolor='none', zorder=zorder)
- ax.add_patch(rect)
- # Draw separator lines between animals
- if idx % 2 == 0:
- ax.axhline(idx - 0.5, color='black', linestyle='-', linewidth=1)
- ytick_color = 'red'
- else:
- ax.axhline(idx - 0.5, color='black', linestyle='--', linewidth=0.1)
- ytick_color = 'green'
- # Set ytick color (handled outside loop, we'll set later)
- # Here only draw lines, not labels
- # Set y-axis labels
- if i == 0:
- # Dynamically generate labels based on animal count (assuming D/S alternating)
- labels = ['D' if j % 2 == 0 else 'S' for j in range(n_animals)]
- ax.set_yticklabels(labels)
- ax.set_ylabel('Rat')
- # Set label colors based on parity
- for j, label in enumerate(ax.get_yticklabels()):
- label.set_color('red' if j % 2 == 0 else 'green')
- else:
- ax.set_yticks([])
- # ========== Step 3: Set x-axis range and right blank rectangle ==========
- ax.set_xticks([])
- # Left white rectangle (per_rect) — keep unchanged!
- per_rect = plt.Rectangle((ret_framerange[0] - ret_mod_width - ret_blank_width, -0.5),
- ret_blank_width, n_animals,
- facecolor='white', edgecolor='white', zorder=8)
- ax.add_patch(per_rect)
- # Right white rectangle (post_rect)
- post_rect = plt.Rectangle((ret_framerange[1], -0.5),
- ret_blank_width, n_animals,
- facecolor='white', edgecolor='white', zorder=8)
- ax.add_patch(post_rect)
- # Set x-axis range: include per_rect, transition, main data, post_rect
- ax.set_xlim(ret_framerange[0] - ret_mod_width - ret_blank_width,
- ret_framerange[1] + ret_blank_width)
- ax.set_ylim(-0.5, n_animals - 0.5)
- ax.invert_yaxis()
- # Add arrow (left of transition zone)
- 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')
- ret_arrow.set_clip_on(False)
- ax.add_patch(ret_arrow)
- # Threat line (using actual x-axis range)
- threat_line = plt.Line2D([0, 8900], [5.75, 5.75], color='red', linewidth=2)
- threat_line.set_clip_on(False)
- ax.add_artist(threat_line)
- # Time scale (only right subplot)
- if i == 1:
- line_x1 = ret_framerange[1] - 30*60
- line_x2 = ret_framerange[1]
- ax.text((line_x1 + line_x2) / 2, -0.8, '60 sec', fontsize=8, color='black', ha='center')
- ret_line = plt.Line2D([line_x1, line_x2], [-0.65, -0.65], color='black', linewidth=1)
- ret_line.set_clip_on(False)
- ax.add_artist(ret_line)
- # Legend remains unchanged
- ret_legend = []
- for label, (color, zorder) in ret_legend_.items():
- if label == 'other behaviors':
- rect = plt.Rectangle((0, 0), 1, 1, facecolor=adjust_saturation(color, saturation),
- edgecolor='black', linewidth=0.5, label=label)
- else:
- rect = plt.Rectangle((0, 0), 1, 1, facecolor=adjust_saturation(color, saturation), label=label)
- ret_legend.append(rect)
- axes[1].legend(handles=ret_legend, loc='upper center', bbox_to_anchor=(0, -0.2), fontsize=7, ncol=7)
- # plt.savefig('fig/fig_2/ret_behavior_ethogram.eps', format="eps", dpi=300, bbox_inches="tight")
- plt.show()
- # %% [markdown]
- # ### Fig2I. Social time(s/(min∙cm2))
- # %%
- import copy
- import numpy as np
- import pandas as pd
- 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
- # Define function to filter social behaviors within far zone
- def filter_social_in_far(social_frames, location_series):
- """
- social_frames: [(start, end), ...]
- location_series: pd.Series, True if in far zone
- Returns: [(far_start, far_end), ...]
- """
- social_in_far = []
- for start, end in social_frames:
- loc_segment = location_series[start:end+1]
- in_far_mask = loc_segment
- if not in_far_mask.any():
- continue
- in_far_indices = loc_segment.index[in_far_mask]
- group_start = None
- for idx in in_far_indices:
- if group_start is None:
- group_start = idx
- prev_idx = idx
- elif idx == prev_idx + 1:
- prev_idx = idx
- else:
- social_in_far.append((group_start, prev_idx))
- group_start = idx
- prev_idx = idx
- if group_start is not None:
- social_in_far.append((group_start, prev_idx))
- if social_in_far == []:
- social_in_far = [np.nan]
- return social_in_far
- # Merge social behavior data
- ap_bl_data = filter_dict_data(ret_origin_bhvr_labels_dict, 'approach_partner')
- sp_bl_data = filter_dict_data(ret_origin_bhvr_labels_dict, 'sniff_partner')
- fp_bl_data = filter_dict_data(ret_origin_bhvr_labels_dict, 'follow_partner')
- gp_bl_data = filter_dict_data(ret_origin_bhvr_labels_dict, 'groom_partner')
- hp_bl_data = filter_dict_data(ret_origin_bhvr_labels_dict, 'huddling')
- merged_bl_data = merge_dicts(ap_bl_data, sp_bl_data)
- merged_bl_data = merge_dicts(merged_bl_data, fp_bl_data)
- merged_bl_data = merge_dicts(merged_bl_data, gp_bl_data)
- merged_bl_data = merge_dicts(merged_bl_data, hp_bl_data)
- ap_wr_data = filter_dict_data(ret_bhvr_labels_dict, 'approach_partner')
- sp_wr_data = filter_dict_data(ret_bhvr_labels_dict, 'sniff_partner')
- fp_wr_data = filter_dict_data(ret_bhvr_labels_dict, 'follow_partner')
- gp_wr_data = filter_dict_data(ret_bhvr_labels_dict, 'groom_partner')
- hp_wr_data = filter_dict_data(ret_bhvr_labels_dict, 'huddling')
- merged_wr_data = merge_dicts(ap_wr_data, sp_wr_data)
- merged_wr_data = merge_dicts(merged_wr_data, fp_wr_data)
- merged_wr_data = merge_dicts(merged_wr_data, gp_wr_data)
- merged_wr_data = merge_dicts(merged_wr_data, hp_wr_data)
- merged_bl_data = filter_in_range(merged_bl_data, [0,8900], method='replace')
- merged_wr_data = filter_in_range(merged_wr_data, [0,8900], method='replace')
- merged_data = {'s': merged_bl_data['p'], 'p': merged_wr_data['p']}
- # Extract position data and filter social behaviors within far zone
- dx = (19 + 5) / 24
- rat_x = -5
- max_frames = 8900
- x_center_al = ret_extract_target_data(ret_raw_data, target_key=['X center'], times=['withrat'])
- social_in_far_data = {}
- for m in ['D', 'S']:
- cond = 'MD' if m == 'D' else 'MS'
- for d in range(1, 12):
- session = f'D{d}{cond}'
- series_x = x_center_al['pair'][cond]['withrat'][f'D{d}']
- s = pd.to_numeric(series_x.iloc[:max_frames], errors='coerce')
- far_mask = ((s >= rat_x + 18 * dx) & (s <= rat_x + 24 * dx)).fillna(False)
- far_series = pd.Series(far_mask.values, index=range(len(far_mask)))
- social_frames = merged_data.get('p', {}).get(m, {}).get(session, [])
- if m not in social_in_far_data:
- social_in_far_data[m] = {}
- if not np.isnan(social_frames).any():
- social_in_far = filter_social_in_far(social_frames, far_series)
- social_in_far_data[m][session] = social_in_far
- else:
- social_in_far_data[m][session] = [np.nan]
- # Calculate duration
- social_in_far_data_dict = {'p': social_in_far_data}
- social_in_far_durations = calculate_duration_time(social_in_far_data_dict)
- total_social_in_far_duration = calculate_total(social_in_far_durations)
- total_social_in_far_duration_df = dict_to_dataframe(total_social_in_far_duration, value_name='in_far_duration', groups=['p'], nan2zero=True)
- raw_social_data = copy.deepcopy(merged_data)
- social_data = {'p': raw_social_data['p']}
- social_durations = calculate_duration_time(social_data)
- total_social_duration = calculate_total(social_durations)
- total_social_duration_df = dict_to_dataframe(total_social_duration, value_name='total_social_duration', groups=['p'], nan2zero=True)
- # Calculate duration in other zones
- total_social_in_other_duration_df = total_social_duration_df[['B_id']].copy()
- for suffix in ['_pD', '_pS']:
- total_col = f'total_social_duration{suffix}'
- in_far_col = f'in_far_duration{suffix}'
- in_other_col = f'in_other_duration{suffix}'
- if total_col in total_social_duration_df.
analyze_data-checkpoint.ipynb at commit 2b40a15, under Apache-2.0 · at the source
Overview
- Department of Neurobiology, School of Basic Medical Sciences, Capital Medical University Beijing China
- Beijing Institute for Brain Research, Chinese Academy of Medical Sciences and Peking Union Medical College Beijing China
- Chinese Institute for Brain Research Beijing China
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
2b40a1596de1298b089c74093646b7e29e2671b7, 13 July 2026Availability: 1 check, the latest on 26 September 2026: the link answers
- 26 September 2026: the link answers
3 files
- .ipynb_checkpoints/
analyze_data-checkpoint. , Jupyter, 5,049 linesipynb - .ipynb_checkpoints/
analyze_data_utils-check , Python, 3,997 linespoint.py - .ipynb_checkpoints/
plot_figures-checkpoint. , Jupyter, 6,041 linesipynb - repository limit reached (2,000 files or 30 MB): the rest is at the source (7 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://
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://
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/
url = {https://
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/
VL - 15
SP - RP109571
SN - 2050-084X
PB - eLife Sciences Publications, Ltd
DO - 10.7554/
UR - https://
LA - en
ER -
CSL-JSON
{
"id": "10.7554/
"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":
"volume": "15",
"page": "RP109571",
"DOI": "10.7554/
"PMID": "42742131",
"PMCID": "PMC13577663",
"ISSN": "2050-084X",
"publisher": "eLife Sciences Publications, Ltd",
"URL": "https://
"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: eLifeIn 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: iScienceIn common: SciPy, Matplotlib, NumPy, mouse, 3 references
- [3] doi:10.7554/elife.109717 [code]
- Retrosplenial cortex enables context-dependent goal-directed sensorimotor transformation.Journal: eLifeIn 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 communicationsIn 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 NeuropsychopharmacologyIn 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 biologyIn 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: iScienceIn 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 researchIn 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 neuroscienceIn 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.
Claim this paper
Correct its record
Say what each link of this record is, remove the ones that are not the paper's, add the ones that are missing. The correction becomes a new version of the record, in its Versions section.
Validate its tracing map
You validate the map as this page shows it: 1 repository of the authors' code, each at its verified commit and with its license, 3 scripts, and 0 matches between paragraphs and code (see the Code and Map sections). It then receives a DOI on Zenodo, with you (your ORCID iD) and OSCR as its creators; the code itself is not deposited.
The map's fingerprint: sha256:1911c50f248a83ab…
Add the badge to its README
The badge links the code to this page. Copy one of these into the README of the paper's code: only you decide where it goes, and nothing is changed for you.
Markdown
[, paste the snippet at the top, then “Commit changes…” and, to review it first, “Create a new branch and start a pull request”. You open the pull request; OSCR asks for no permission.
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.
