{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":106809,"databundleVersionId":13056355,"sourceType":"competition"},{"sourceId":13306764,"sourceType":"datasetVersion","datasetId":8434700}],"dockerImageVersionId":31153,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport h5py\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom tqdm.notebook import tqdm\n\n# --- Style and Configuration ---\nplt.style.use('seaborn-v0_8-whitegrid')\nsns.set_context(\"talk\") # 'talk' context for larger, more readable plots\npd.set_option('display.max_colwidth', 100) # Show more of the sentence text\n\n# --- Constants ---\nBASE_DIR = '/kaggle/input/brain-to-text-25/t15_copyTask_neuralData/hdf5_data_final'\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Loading and Structuring the Data\n\nThe data is organized into sessions by date, with train, val, test splits in HDF5 files. \nLet's load the metadata (everything ex except the large neural data arrays) into a single pandas DataFrame for easy analysis","metadata":{}},{"cell_type":"code","source":"import h5py\nimport numpy as np\nimport pandas as pd\nfrom tqdm.notebook import tqdm\nimport os\nfrom datetime import datetime\nimport pyarrow as pa\n\nNEURAL_DATA_KEY = 'input_features'\nTRANSCRIPTION_KEY = 'transcription'\n\n\ndef load_metadata_from_hdf5(file_path , df_description):\n    metadata = []\n    try:\n        with h5py.File(file_path, 'r') as f:\n            trials = list(f.keys()) # for each trial\n            for trial in trials:\n                \n                dataset_names = list(f[trial].keys()) # ['input_features', 'seq_class_ids', 'transcription']\n                # check if the group contains the correct dataset names\n                if isinstance(f[trial], h5py.Group) and NEURAL_DATA_KEY in dataset_names and TRANSCRIPTION_KEY in dataset_names:\n                    data_attr = f[trial].attrs.keys() # ['block_num', 'n_time_steps', 'sentence_label', 'seq_len', 'session', 'trial_num']\n                    \n                    block_num = f[trial].attrs['block_num'] \n                    sentence_label = f[trial].attrs['sentence_label'] \n                    phoneme_seq_len = f[trial].attrs['seq_len'] \n                    session = f[trial].attrs['session'] \n                    trial_num = f[trial].attrs['trial_num'] \n\n                    # number of words in transcription\n                    num_words = len(sentence_label.split())\n\n                    input_features = f[trial]['input_features'][()]\n      \n                    num_time_bins, num_channels = input_features.shape\n                    # [()] on an HDF5 dataset in h5py is a shorthand for “read the entire dataset into memory as a NumPy array\n                    seq_class_ids = f[trial]['seq_class_ids'][()] \n                    seq_transcription = list(f[trial]['transcription'][()])\n\n                    # reformat the session into proper date format\n                    date_part = session.split('.', 1)[1]\n                    dt = datetime.strptime(date_part, \"%Y.%m.%d\")\n                    formatted_date = dt.strftime(\"%Y-%m-%d\")\n\n                    corpus_list = df_description[(df_description['Date'] == formatted_date) & (df_description['Block number'] == block_num)\n                          ]['Corpus'].values\n\n                    if len(corpus_list) > 0: # since it is inside an array\n                        corpus = corpus_list[0]\n                    else:\n                        corpus = None\n\n                    metadata.append({\n                        \n                        'session': session,# session e.g. t15.2023.08.11...\n                        \n                        'trial_id': trial, # current trial_id, represents the cumulated trials within current session\n                        \n                        'block_number' : block_num, # Research block number that the trial is sourced from\n                        \n                        'trial_num': trial_num, # trial_num: depends on the number of sentences within this block, if number of sentence = 20, trial_num = 0-19\n\n                        'corpus': corpus, # [Switchboard, Harvard sentences, OpenWebText2, Custom high-frequency word sentences, Random word sentences]\n                        \n                        'num_time_bins': num_time_bins, # Number of time steps per trial\n                        # neural features\n                        'neural_features': input_features,\n                        # Phonemes\n                        'phoneme_labels' : seq_class_ids, #  Integer phoneme sequence labels for each trial. [ 7 28 17 24 40 17 31 40 20 21 25 29 12 40  0  0...]\n                        'num_of_phoneme_labels' : phoneme_seq_len, # Number of phoneme labels per trial.\n                        # Transcriptions\n                        'transcription_ASCII_characters': seq_transcription, # ASCII representation of sentence label for each trial. [ 66 114 105 110 103  32 105 116  32  99 108 111 115 101 114  46   0   0...]\n                        'num_of_ASCII_characters': len(sentence_label),\n                        \n                        'transcription_text': sentence_label,\n                        'num_texts' : num_words,\n            \n                    }) \n                    \n            \n    except Exception as e:\n        print(f\"Error processing file {file_path}: {e}\")\n        import traceback\n        traceback.print_exc()\n        \n    return metadata\n    \ndef load_test_metadata_from_hdf5(file_path, df_description):\n    metadata = [] # a list of dicts\n    try:\n        with h5py.File(file_path, 'r') as f:\n            trials = list(f.keys()) # for each trial\n            for trial in trials:\n                \n                dataset_names = list(f[trial].keys()) # ['input_features']\n\n                # check if the group contains the correct dataset names\n                if isinstance(f[trial], h5py.Group) and NEURAL_DATA_KEY in dataset_names: # test data which only has input_features. \n                    data_attr = f[trial].attrs.keys() # ['block_num', 'n_time_steps', 'session', 'trial_num'] # doesn't have sentence_label and seq_len\n\n                    block_num = f[trial].attrs['block_num'] \n                    session = f[trial].attrs['session'] \n                    trial_num = f[trial].attrs['trial_num'] \n\n                    input_features = f[trial]['input_features'][()]\n         \n                    num_time_bins, num_channels = input_features.shape\n\n                    # reformat the session into proper date format\n                    date_part = session.split('.', 1)[1]\n                    dt = datetime.strptime(date_part, \"%Y.%m.%d\")\n                    formatted_date = dt.strftime(\"%Y-%m-%d\")\n\n                    corpus_list = df_description[(df_description['Date'] == formatted_date) & (df_description['Block number'] == block_num)]['Corpus'].values\n\n                    if len(corpus_list) > 0: # since it is inside an array\n                        corpus = corpus_list[0]\n                    else:\n                        corpus = None\n                    \n                        \n                    metadata.append({\n                        \n                        'session': session,# session e.g. t15.2023.08.11...\n                        \n                        'trial_id': trial, # current trial_id, represents the cumulated trials within current session\n                        \n                        'block_number' : block_num, # Research block number that the trial is sourced from\n                        \n                        'trial_num': trial_num, # trial_num: depends on the number of sentences within this block, if number of sentence = 20, trial_num = 0-19\n\n                        'corpus': corpus,\n                        \n                        'num_time_bins': num_time_bins, # Number of time steps per trial\n                        # neural features\n                        'neural_features': input_features,\n\n                    }) \n           \n    except Exception as e:\n        print(f\"Error processing file {file_path}: {e}\")\n        import traceback\n        traceback.print_exc()\n    return metadata\n\n# -- Main Loading Loop --\n\nall_metadata = []\n\nsession_dirs = sorted([d for d in os.listdir(BASE_DIR) if os.path.isdir(os.path.join(BASE_DIR, d))])\n\nt15_copyTaskData_description_path = '/kaggle/input/brain-to-text-25-copytaskdata-description/t15_copyTaskData_description.csv'\ndf_description = pd.read_csv(t15_copyTaskData_description_path)\n\nfor session in tqdm(session_dirs, desc='Processing Sessions'):\n    session_path = os.path.join(BASE_DIR, session)\n    for split in ['train', 'val', 'test']:\n        file_name = f'data_{split}.hdf5'\n        file_path = os.path.join(session_path, file_name)\n\n        if os.path.exists(file_path):\n            if split == 'test':\n                session_metadata = load_test_metadata_from_hdf5(file_path, df_description)\n            else:\n                session_metadata = load_metadata_from_hdf5(file_path, df_description)\n                \n            # Add session and split info to the found trials\n            for item in session_metadata: # for every dict in the list\n                item['split'] = split # assign train/test/split\n\n            all_metadata.extend(session_metadata)\n\n\ndf = pd.DataFrame(all_metadata)\nprint(df.dtypes)           \nprint(f\"Loaded a total of {len(df)} trials. \")\nif not df.empty:\n    print(f\"Data splits: \\n{df['split'].value_counts()}\")\n    display(df.head())\n    df.to_pickle('metadata.pkl')\n    # df.to_parquet('metadata.parquet', index=False)\n    # df.to_csv('/kaggle/working/metadata.csv', index=False)\nelse:\n    print(\"DataFrame is still empty. This indicates a very unusual issue.\")   \n            ","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport sys\n\n# Helper function to format bytes into a readable string (KB, MB, GB)\ndef format_bytes(byte_count):\n    if byte_count is None:\n        return \"N/A\"\n    power = 1024\n    n = 0\n    power_labels = {0: '', 1: 'K', 2: 'M', 3: 'G', 4: 'T'}\n    while byte_count >= power and n < len(power_labels):\n        byte_count /= power\n        n += 1\n    return f\"{byte_count:.2f} {power_labels[n]}B\"\n\n\n# --- Method 1: df.info() ---\n# Good for a quick, high-level summary.\nprint(\"--- Method 1: Using df.info() ---\")\n\nprint(\"\\nShallow memory usage (default):\")\ndf.info(verbose=False) # verbose=False keeps the output tidy\n\nprint(\"\\nDeep memory usage (more accurate):\")\ndf.info(verbose=False, memory_usage='deep')\n\n\n# --- Method 2: df.memory_usage() ---\n# Good for getting the raw numbers and seeing memory per column.\nprint(\"\\n\\n--- Method 2: Using df.memory_usage() ---\")\n\n# Shallow memory calculation\nshallow_mem = df.memory_usage(index=True, deep=False).sum()\nprint(f\"Shallow memory total: {shallow_mem} bytes ({format_bytes(shallow_mem)})\")\n\n# Deep memory calculation (the one you usually want)\ndeep_mem = df.memory_usage(index=True, deep=True).sum()\nprint(f\"Deep memory total:    {deep_mem} bytes ({format_bytes(deep_mem)})\")\n\nprint(\"\\nBreakdown by column (Deep):\")\nprint(df.memory_usage(index=True, deep=True).apply(format_bytes))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Check NaN ( all NaN values are test data)","metadata":{}},{"cell_type":"code","source":"# 1. Check which columns contain any NaNs (returns a Boolean per column)\nnan_presence = df.isna().any()\nprint(nan_presence)\nprint()\n# 2. Count the number of NaNs in each column (returns an integer count per column)\nnan_counts = df.isna().sum()\nprint(nan_counts)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Check for duplicate values\n\nNote that within the same day, train/val/test data shares the same session, and the trial_id counts from 'trial_0000', 'trial_0001' for train/val/test respectively, so we get duplicated values\n\ntranscription text shows no duplicates, indicating no training data are being repeated","metadata":{}},{"cell_type":"code","source":"# Boolean mask: True for rows where (session, trial_id) duplicates a previous row\ndup_mask = df.duplicated(subset=['session', 'trial_id'])\n\n# Print mask\nprint(\"Duplicate (session, trial_id) mask:\")\nprint(dup_mask)\n\n# Count how many duplicate rows there are\nnum_dups = dup_mask.sum()\nprint(f\"\\nNumber of duplicate (session, trial_id) entries: {num_dups}\")\n\n# To see all rows that share a duplicate key (marking every copy)\nall_dups = df[df.duplicated(subset=['session', 'trial_id'], keep=False)][['session', 'trial_id', 'transcription_text']]\nprint(\"\\nRows with duplicate (session, trial_id):\")\nprint(all_dups)\n\nprint(\"Exported duplicates to session_trial_duplicates.csv\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Check duplicates based on session, trial_id, and transcription_text\n\n# 1. Boolean mask: True for rows where (session, trial_id, transcription_text) duplicates a previous row\ndup_mask = df.duplicated(subset=['session', 'trial_id', 'transcription_text'])\n\n# 2. Print mask\nprint(\"Duplicate (session, trial_id, transcription_text) mask:\")\nprint(dup_mask)\n\n# 3. Count how many duplicate rows there are\nnum_dups = dup_mask.sum()\nprint(f\"\\nNumber of duplicate (session, trial_id, transcription_text) entries: {num_dups}\")\n\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Sentence Length Distribution","metadata":{}},{"cell_type":"code","source":"plt.figure(figsize=(12, 6))\n# We filter where num_words is not NaN to exclude the test set\nsns.histplot(data=df.dropna(subset=['num_texts']), x='num_texts', bins=np.arange(0, 40, 1), kde=False)\nplt.title('Distribution of Sentence Length (Number of Texts)')\nplt.xlabel('Number of Texts')\nplt.ylabel('Count')\nplt.show()\n\nprint(\"Sentence Length Statistics:\")\nprint(df['num_texts'].describe())","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Vocabulary Size\nif we don't get rid of punctuations:\nTotal number of words (tokens): 60094\nVocabulary size (unique words): 5979\n\nAfter we remove punctuations:\nTotal number of words (tokens): 60091\nVocabulary size (unique words): 4279","metadata":{}},{"cell_type":"code","source":"import string\n# We need to tokenize the sentence to count unique words\n# ' '.join(...) → Combines all those sentences into one long string, separated by spaces\n# string.punctuation contains all standard punctuation: !\"#$%&'()*+,-./:;<=>?@[$$^_{|}~`\nall_words = ' '.join(df[df['split'] != 'test']['transcription_text']).translate(str.maketrans('', '', string.punctuation)).lower().split()\nvocabulary = set(all_words)\n\nprint(f\"Total number of words (tokens): {len(all_words)}\")\nprint(f\"Vocabulary size (unique words): {len(vocabulary)}\")\nprint(\"\\nSome example words from the vocabulary:\")\n# Filter out any empty strings that might result from splitting\nprint([word for word in list(vocabulary) if word][:15])","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Trial Duration and Speaking Rate","metadata":{}},{"cell_type":"code","source":"# Duration in seconds = num_time_bins * 0.02, since each bin is sampled at 20ms = 0.02s\ndf['duration_s'] = df['num_time_bins'] * 0.02\n\n# WPM -- Word-Per-Minute\ndf['wpm'] = df['num_texts'] / df['duration_s'] * 60\n\nplt.figure(figsize=(18, 6))\n\n# Plot 1: Duration\nplt.subplot(1, 2, 1)\nsns.histplot(data=df, x='duration_s', bins=50)\nplt.title('Distribution of Trial Durations')\nplt.xlabel('Duration (seconds)')\nplt.ylabel('Count')\n\n# Plot 2: Words Per Minute\nplt.subplot(1, 2, 2)\n# Filter out NaNs from the test set for the WPM plot\nsns.histplot(data=df.dropna(subset=['wpm']), x='wpm', bins=50)\nplt.title('Distribution of Speaking Rate (WPM)')\nplt.xlabel('Words Per Minute (WPM)')\nplt.ylabel('Count')\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Relationship between Recording Length and phoneme lengths","metadata":{}},{"cell_type":"code","source":"print(df[df['num_time_bins'] == df['num_time_bins'].max()])","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(df[df['num_of_phoneme_labels'] == df['num_of_phoneme_labels'].max()])","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Combine train and test with available transcription lengths\nmask = (df['split'] == 'train') | (df['split'] == 'val')\nsubset = df[mask].copy()\n\nplt.figure(figsize=(6,6))\nsns.scatterplot(x='num_time_bins', y='num_of_phoneme_labels', data=subset)\nplt.title(\"Time Bins vs. Phoneme Labels\")\nplt.xlabel(\"Neural Time Bins\")\nplt.ylabel(\"Number of Phoneme labels\")\nplt.show()\n\n# Compute correlation\ncorr = subset[['num_time_bins','num_of_phoneme_labels']].corr().iloc[0,1]\nprint(f\"Correlation: {corr:.2f}\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom scipy import stats\n\n# Your existing code for scatter plot...\n# mask = (df['split'] == 'train') | (df['split'] == 'val')\n# subset = df[mask].copy()\n\n# ADD THIS: Calculate the ratio\nsubset['ratio'] = subset['num_time_bins'] / subset['num_of_phoneme_labels']\nratio = subset['ratio']\n\n# Comprehensive statistics\nprint(\"\\n\" + \"=\"*60)\nprint(\"RATIO STATISTICS (Neural Time Bins / Phoneme Labels)\")\nprint(\"=\"*60)\nprint(f\"Count: {len(ratio):,}\")\nprint(f\"Min ratio: {ratio.min():.4f}\")\nprint(f\"Max ratio: {ratio.max():.4f}\")\nprint(f\"Mean ratio: {ratio.mean():.4f}\")\nprint(f\"Median ratio: {ratio.median():.4f}\")\nprint(f\"Standard deviation: {ratio.std():.4f}\")\n\nprint(f\"\\nADDITIONAL STATISTICS:\")\nprint(f\"25th percentile: {ratio.quantile(0.25):.4f}\")\nprint(f\"75th percentile: {ratio.quantile(0.75):.4f}\")\nprint(f\"IQR: {ratio.quantile(0.75) - ratio.quantile(0.25):.4f}\")\nprint(f\"CV: {(ratio.std() / ratio.mean()) * 100:.2f}%\")\nprint(f\"Skewness: {stats.skew(ratio):.4f}\")\nprint(f\"Kurtosis: {stats.kurtosis(ratio):.4f}\")\n\n# Percentile analysis\npercentiles = [1, 5, 10, 25, 50, 75, 90, 95, 99]\nprint(f\"\\nPERCENTILE ANALYSIS:\")\nfor p in percentiles:\n    print(f\"{p:2d}th percentile: {np.percentile(ratio, p):8.4f}\")\n\n# Create ratio distribution plots\nplt.figure(figsize=(15, 5))\n\n# Histogram\nplt.subplot(1, 3, 1)\nplt.hist(ratio, bins=50, alpha=0.7, color='skyblue', edgecolor='black')\nplt.axvline(ratio.mean(), color='red', linestyle='--', label=f'Mean: {ratio.mean():.2f}')\nplt.axvline(ratio.median(), color='green', linestyle='--', label=f'Median: {ratio.median():.2f}')\nplt.xlabel('Ratio (Time Bins / Phoneme Labels)')\nplt.ylabel('Frequency')\nplt.title('Distribution of Ratios')\nplt.legend()\nplt.grid(alpha=0.3)\n\n# Box plot\nplt.subplot(1, 3, 2)\nplt.boxplot(ratio)\nplt.ylabel('Ratio Value')\nplt.title('Box Plot of Ratios')\nplt.grid(alpha=0.3)\n\n# Q-Q plot\nplt.subplot(1, 3, 3)\nstats.probplot(ratio, dist=\"norm\", plot=plt)\nplt.title('Q-Q Plot vs Normal')\nplt.grid(alpha=0.3)\n\nplt.tight_layout()\nplt.show()\n\n# Summary\nprint(f\"\\nSUMMARY:\")\nprint(f\"On average, there are {ratio.mean():.2f} neural time bins per phoneme label.\")\nprint(f\"50% of data has ratios between {ratio.quantile(0.25):.2f} and {ratio.quantile(0.75):.2f}.\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Corpus Distribution","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport seaborn as sns\n\nsplits = ['train', 'val', 'test']\nfig, axes = plt.subplots(1, 3, figsize=(15, 5), sharey=True)\n\nfor ax, split in zip(axes, splits):\n    sub = df[df['split'] == split]\n    order = sub['corpus'].value_counts().index\n    total = len(sub)\n\n    # Countplot\n    sns.countplot(\n        y='corpus',\n        data=sub,\n        order=order,\n        palette='pastel',\n        ax=ax\n    )\n    ax.set_title(f\"{split.capitalize()} Split\\n(n={total})\")\n    ax.set_xlabel(\"Trial Count\")\n    ax.set_ylabel(\"\" if split != 'train' else \"Corpus\")\n\n    # Annotate percentages\n    for p in ax.patches:\n        count = p.get_width()\n        percentage = 100 * count / total\n        ax.text(\n            count + total * 0.005,  # slight offset\n            p.get_y() + p.get_height() / 2,\n            f\"{percentage:.1f}%\",\n            va='center'\n        )\n\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Group by corpus and compute means\ncorpus_summary = df.groupby('corpus').agg(\n    mean_time_bins = ('num_time_bins', 'mean'),\n    mean_phoneme_labels = ('num_of_phoneme_labels', 'mean'),\n    mean_duration_s = ('duration_s', 'mean'),\n    mean_wpm       = ('wpm', 'mean'),\n    trial_count    = ('trial_id', 'count')\n).reset_index()\n\n# Assume corpus_summary is already defined\norder = corpus_summary.sort_values('trial_count', ascending=False)['corpus']\n\nmetrics = [\n    ('mean_time_bins',       'Average Time Bins'),\n    ('mean_phoneme_labels',  'Average Phoneme Count'),\n    ('mean_duration_s',      'Average Duration (s)'),\n    ('mean_wpm',             'Average Words per Minute')\n]\n\nfig, axes = plt.subplots(2, 2, figsize=(12, 10))\naxes = axes.flatten()\n\nfor ax, (col, title) in zip(axes, metrics):\n    sns.barplot(\n        x='corpus',\n        y=col,\n        data=corpus_summary,\n        order=order,\n        palette='viridis',\n        ax=ax\n    )\n    ax.set_title(f'{title} by Corpus')\n    ax.set_xlabel('')\n    ax.set_ylabel(title)\n    ax.tick_params(axis='x', rotation=45)\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Pearson correlation matrices between channel pairs","metadata":{}},{"cell_type":"code","source":"# --- Define feature mappings (this part is still correct) ---\nFEATURE_MAP = {\n    'ventral_6v_thresh': list(range(0, 64)),\n    'area_4_thresh': list(range(64, 128)),\n    'area_55b_thresh': list(range(128, 192)),\n    'dorsal_6v_thresh': list(range(192, 256)),\n    'ventral_6v_sbp': list(range(256, 320)),\n    'area_4_sbp': list(range(320, 384)),\n    'area_55b_sbp': list(range(384, 448)),\n    'dorsal_6v_sbp': list(range(448, 512)),\n}\n\n# Define the 4 threshold regions\nFEATURE_MAP_THRESH = {\n    'ventral_6v': FEATURE_MAP['ventral_6v_thresh'],\n    'area_4':    FEATURE_MAP['area_4_thresh'],\n    'area_55b':  FEATURE_MAP['area_55b_thresh'],\n    'dorsal_6v': FEATURE_MAP['dorsal_6v_thresh'],\n}\n\nneural_features_list = df[df['split'] != 'test']['neural_features'].tolist() # neural_features_list is a list of all data (length 9498), each element inside represents 1 numpy array of shape (num_timestep, num_channels)\n\n# Mean across time_bins \nchannel_means = np.array([trial.mean(axis=0) for trial in neural_features_list]) # aggregate mean across num_timsteps, which gives us a remaining 2-D numpy array of shape (all_data, 512)\n\n# 2. Select only the first 256 channels (threshold features)\nchannel_means_threshold = channel_means[:, :256]\nchannel_means_spike_band_power = channel_means[:, 256:]\n\"\"\"\n# Compute region averages\nregion_means = np.zeros((channel_means.shape[0], len(FEATURE_MAP_THRESH))) # (all data, 4 regions)\nfor idx, (region, ch_idx) in enumerate(FEATURE_MAP_THRESH.items()):\n    # take the mean across region channels\n    region_means[:, idx] = channel_means[:, ch_idx].mean(axis=1) # region_means[all_data, index = current_region] = channel_means[all_data, 0-64]\n\n\"\"\"\n# 3. Compute the 256×256 correlation matrix across channels\n#    rowvar=False ensures we correlate columns (channels) across trials\ncorr_matrix = np.corrcoef(channel_means_spike_band_power, rowvar=False)\n\n# 4. Plot the correlation matrix\nfig, ax = plt.subplots(figsize=(8, 6))\nim = ax.imshow(corr_matrix, cmap='RdBu_r', vmin=-1, vmax=1, aspect='auto')\nplt.colorbar(im, ax=ax, fraction=0.046, pad=0.04)\n\n# 5. Label axes (optional: provide your channel names or indices)\nchannel_indices = np.arange(256)\nax.set_xticks(channel_indices[::16])           # e.g., label every 16th channel\nax.set_yticks(channel_indices[::16])\nax.set_xticklabels(channel_indices[::16], rotation=45)\nax.set_yticklabels(channel_indices[::16])\nax.set_xlabel('Channel Index')\nax.set_ylabel('Channel Index')\nax.set_title('Correlation Matrix of Last 256 Channels (spike_band_power Features)')\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Time-Series Analysis","metadata":{}},{"cell_type":"markdown","source":"### Neural Features for a single trial","metadata":{}},{"cell_type":"code","source":"def plot_single_trial(session, trial_id):\n    '''\n    Loads and plots a single trial's neural data and sentence,\n    optimizing for visibility of neural activity spikes\n    '''\n    # Find the correct file and load the data (no changes here)\n    row = df[(df['session'] == session) & (df['trial_id'] == trial_id)].iloc[0]\n    split = row['split']\n    file_path = os.path.join(BASE_DIR, session, f'data_{split}.hdf5')\n    \n    NEURAL_DATA_KEY = 'input_features'\n    \n    with h5py.File(file_path, 'r') as f:\n        trial_group = f[trial_id]\n        neural_data = trial_group[NEURAL_DATA_KEY][()]\n        sentence = row['transcription_text']\n\n    # --- Plotting Enhancements ---\n    \n    # 1. Create a proper time axis in seconds for the x-axis\n    time_in_seconds = np.arange(neural_data.shape[0]) * 0.02\n    \n    plt.figure(figsize=(20, 10))\n\n    # 2. Use imshow with parameters for better visibility\n    plt.imshow(\n        neural_data.T,\n        aspect=0.02,  # Stretch the plot horizontally to make spikes wider\n        cmap='magma', # A colormap with high contrast (dark background, bright spikes)\n        interpolation='none',\n        vmin=0,       # Set the minimum color value to 0 to ignore negative noise\n        vmax=8,       # Clip the max color value to make spikes pop \n        extent=[time_in_seconds[0], time_in_seconds[-1], neural_data.shape[1], 0] # Correctly scale axes\n    )\n\n    # Add lines to separate the brain regions\n    for i in [64, 128, 192, 256, 320, 384, 448]:\n        plt.axhline(y=i-0.5, color='white', linestyle='--', linewidth=1)\n        \n    plt.colorbar(label='Neural Activity (Clipped at 8)')\n    plt.yticks(\n        ticks=[32, 96, 160, 224, 288, 352, 416, 480],\n        labels=[\n            'v6v (Thresh)', 'Area 4 (Thresh)', '55b (Thresh)', 'd6v (Thresh)',\n            'v6v (SBP)', 'Area 4 (SBP)', '55b (SBP)', 'd6v (SBP)'\n        ],\n        rotation=0\n    )\n    plt.xlabel('Time (seconds)') # labeled in seconds\n    plt.ylabel('Brain Region & Feature Type')\n    plt.title(f'\"{sentence}\"', fontsize=24, pad=20)\n    plt.show()\n\n# --- Call the function ---\n# Pick a trial to visualize\nexample_trial_info = df.dropna(subset=['num_texts'])[df['num_texts'] > 5].iloc[0]\nplot_single_trial(example_trial_info['session'], example_trial_info['trial_id'])","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(df['neural_features'][0][:, 0])","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(df['neural_features'][0][:, 256])","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}