{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"colab":{"provenance":[]},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","sourceId":59093,"databundleVersionId":7469972},{"sourceType":"modelInstanceVersion","sourceId":243694,"databundleVersionId":10934234,"modelInstanceId":208170,"modelId":229872},{"sourceType":"modelInstanceVersion","sourceId":243711,"databundleVersionId":10934368,"modelInstanceId":208185,"modelId":229884}],"dockerImageVersionId":30840,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"! pip install focal-loss-torch","metadata":{"execution":{"iopub.status.busy":"2025-01-28T02:34:32.308472Z","iopub.execute_input":"2025-01-28T02:34:32.308765Z","iopub.status.idle":"2025-01-28T02:34:37.59985Z","shell.execute_reply.started":"2025-01-28T02:34:32.308731Z","shell.execute_reply":"2025-01-28T02:34:37.59872Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# HMS - Harmful Brain Activity Classification","metadata":{"id":"3_jzt7rCE5PA"}},{"cell_type":"markdown","source":"TODOS\n\n1. Graph of uniques sum_of #votes against cont of #eeg_samples  ## Done\n2. confidence vs #samples eg 100% -- no of samples with 100% votes ## Done\n3. do your own inference about this data and share plots \n4. make a csv of eeg_id,eeg_sub_it for 100% sure class consensus\n5, come up with your idea for this\n","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport time\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom tqdm import tqdm\nfrom tqdm.notebook import tqdm as jptq\nimport sys\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.nn import CrossEntropyLoss\nimport torch.utils.data as data\nfrom sklearn.model_selection import train_test_split\nimport numpy as np\nfrom torch.utils.data import Dataset, TensorDataset, DataLoader\nimport torch.optim as optim\nfrom bisect import bisect_left\nfrom focal_loss.focal_loss import FocalLoss\n\nimport warnings\nimport pickle\nimport subprocess\nimport traceback\nfrom concurrent.futures import ProcessPoolExecutor, as_completed, ThreadPoolExecutor\nimport os\nwarnings.filterwarnings('ignore')\n\nfrom transformers import Wav2Vec2ForSequenceClassification, Wav2Vec2Processor\nfrom sklearn.preprocessing import StandardScaler\nfrom scipy.signal import resample\nfrom sklearn.metrics import accuracy_score, f1_score, precision_score, recall_score\n\nfrom sklearn.metrics import accuracy_score, f1_score, precision_score, recall_score\nfrom sklearn.metrics import roc_auc_score, precision_recall_curve, cohen_kappa_score, classification_report\n\nfrom IPython.display import FileLink, display","metadata":{"id":"a6684e18","execution":{"iopub.status.busy":"2025-01-28T02:34:37.600981Z","iopub.execute_input":"2025-01-28T02:34:37.601346Z","iopub.status.idle":"2025-01-28T02:35:00.785804Z","shell.execute_reply.started":"2025-01-28T02:34:37.601309Z","shell.execute_reply":"2025-01-28T02:35:00.785013Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## LOADING THE DATASET","metadata":{"id":"Jeaws__qE5PH"}},{"cell_type":"code","source":"model_folder_path=\"/kaggle/input/model_0/transformers/default/2\"\nmodel_path=\"/kaggle/input/model_0/transformers/default/2/model_1.pkl\"\ndataset_folder_path=\"/kaggle/input/model_0/transformers/default/2\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-28T02:35:00.786716Z","iopub.execute_input":"2025-01-28T02:35:00.7873Z","iopub.status.idle":"2025-01-28T02:35:00.791363Z","shell.execute_reply.started":"2025-01-28T02:35:00.787272Z","shell.execute_reply":"2025-01-28T02:35:00.790466Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"BASE_DIR = \"/kaggle/input/hms-harmful-brain-activity-classification/\"\ndf = pd.read_csv(f\"{BASE_DIR}train.csv\")\ndf['total_votes'] = df['seizure_vote']+df['lpd_vote']+df['gpd_vote']+df['lrda_vote']+df['grda_vote']+df['other_vote']\ndf['seizure_vote_normed'] = df['seizure_vote']/df['total_votes']\ndf['lpd_vote_normed'] = df['lpd_vote']/df['total_votes']\ndf['gpd_vote_normed'] = df['gpd_vote']/df['total_votes']\ndf['lrda_vote_normed'] = df['lrda_vote']/df['total_votes']\ndf['grda_vote_normed'] = df['grda_vote']/df['total_votes']\ndf['other_vote_normed'] = df['other_vote']/df['total_votes']\n\ntarget_encoding = {'Seizure':0, 'LPD':1, 'GPD':2, 'LRDA':3, 'GRDA':4, 'Other':5,\n                   'seizure':0, 'lpd':1, 'gpd':2, 'lrda':3, 'grda':4, 'other':5,\n                   0:'seizure', 1:'lpd', 2:'gpd', 3:'lrda', 4:'grda', 5:'other'}\ntotal_targets = 6\ndf.head()","metadata":{"id":"ab2bfbe1","outputId":"3e9fd5d9-f66a-4ead-e3a3-676461244431","execution":{"iopub.status.busy":"2025-01-28T02:35:00.79251Z","iopub.execute_input":"2025-01-28T02:35:00.7929Z","iopub.status.idle":"2025-01-28T02:35:01.090441Z","shell.execute_reply.started":"2025-01-28T02:35:00.792853Z","shell.execute_reply":"2025-01-28T02:35:01.0896Z"},"trusted":true,"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## DATA ANALYSIS\n\n### Dataframe attributes and Important results\n1. **`shape`**: (106800, 15)\n2. `df.isnull().any() = False` (for all columns)\n\n### Columns:\n\n1. **`eeg_id`**: Unique identifier for the entire EEG recording. It represents a specific EEG session or recording from a patient.\n\n2. **`eeg_sub_id`**: An ID for the specific 50-second long subsample to which the row's labels apply. This column helps identify a particular subset or segment within the larger EEG recording.\n\n3. **`eeg_label_offset_seconds`**: The time between the beginning of the consolidated EEG and the subsample. It indicates the time offset for the EEG subsample within the entire EEG recording.\n\n4. **`spectrogram_id`**: Unique identifier for the entire EEG recording, similar to `eeg_id`. It is related to the spectrogram data.\n\n5. **`spectrogram_sub_id`**: An ID for the specific 10-minute subsample to which the row's labels apply. This corresponds to a subset within the larger spectrogram data.\n\n6. **`spectrogram_label_offset_seconds`**: The time between the beginning of the consolidated spectrogram and the subsample. It indicates the time offset for the spectrogram subsample within the entire spectrogram recording.\n\n7. **`label_id`**: An ID for this set of labels. It helps distinguish different sets of labels within the dataset.\n\n8. **`patient_id`**: An ID for the patient who donated the data. It uniquely identifies each patient.\n\n9. **`expert_consensus`**: The consensus annotator label for convenience. This column may provide a summary or agreement among expert annotators regarding the type of brain activity in the given subsample.\n\n10. **`seizure_vote`**, **`lpd_vote`**, **`gpd_vote`**, **`lrda_vote`**, **`grda_vote`**, **`other_vote`**: These columns represent the count of annotator votes for specific brain activity classes. The classes are:\n    - `seizure_vote`: Count of votes for seizure.\n    - `lpd_vote`: Count of votes for lateralized periodic discharges.\n    - `gpd_vote`: Count of votes for generalized periodic discharges.\n    - `lrda_vote`: Count of votes for lateralized rhythmic delta activity.\n    - `grda_vote`: Count of votes for generalized rhythmic delta activity.\n    - `other_vote`: Count of votes for other types of brain activity.\n\n11. **Target Variable**: The target variable in this dataset is the actual brain activity class for each subsample. It could be any of the following:\n    - Seizure (`seizure_vote`): Represents the count of votes for seizure.\n    - Lateralized Periodic Discharges (`lpd_vote`): Represents the count of votes for lateralized periodic discharges.\n    - Generalized Periodic Discharges (`gpd_vote`): Represents the count of votes for generalized periodic discharges.\n    - Lateralized Rhythmic Delta Activity (`lrda_vote`): Represents the count of votes for lateralized rhythmic delta activity.\n    - Generalized Rhythmic Delta Activity (`grda_vote`): Represents the count of votes for generalized rhythmic delta activity.\n    - Other (`other_vote`): Represents the count of votes for other types of brain activity.\n\nThe target variable is the type of brain activity class (seizure, lpd, gpd, lrda, grda, other), and the goal of the competition or analysis is likely to predict or classify the correct brain activity class for each EEG subsample.","metadata":{"id":"24fe4409"}},{"cell_type":"code","source":"object_columns = df.select_dtypes(include=['object', 'bool']).columns\n# print(\"Object type columns:\")\n# print(object_columns)\nnumerical_columns = df.select_dtypes(include=['int64', 'float64']).columns\n# print(\"\\nNumerical type columns:\")\n# print(numerical_columns)","metadata":{"execution":{"iopub.status.busy":"2025-01-28T02:35:01.091249Z","iopub.execute_input":"2025-01-28T02:35:01.091482Z","iopub.status.idle":"2025-01-28T02:35:01.122976Z","shell.execute_reply.started":"2025-01-28T02:35:01.091462Z","shell.execute_reply":"2025-01-28T02:35:01.121809Z"},"trusted":true,"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Datatype analysis","metadata":{}},{"cell_type":"code","source":"# for col in df.columns:\n#     print(col)\n#     if not isinstance(df[col][0], str):\n#         for i,el in enumerate(tqdm(df[col])):\n#             if int(el)!=el:\n#                 print(i, el)","metadata":{"execution":{"iopub.status.busy":"2025-01-28T02:35:01.125371Z","iopub.execute_input":"2025-01-28T02:35:01.125611Z","iopub.status.idle":"2025-01-28T02:35:01.129228Z","shell.execute_reply.started":"2025-01-28T02:35:01.125591Z","shell.execute_reply":"2025-01-28T02:35:01.128231Z"},"trusted":true,"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Seen\ndef classify_features(df):\n    categorical_features = []\n    non_categorical_features = []\n    discrete_features = []\n    continuous_features = []\n    for column in df.columns:\n        if df[column].dtype in ['object', 'bool']:\n            if df[column].nunique() < 15:\n                categorical_features.append(column)\n            else:\n                non_categorical_features.append(column)\n        elif df[column].dtype in ['int64', 'float64']:\n            if df[column].nunique() < 10:\n                discrete_features.append(column)\n            else:\n                continuous_features.append(column)\n    return categorical_features, non_categorical_features, discrete_features, continuous_features\n\ncategorical, non_categorical, discrete, continuous = classify_features(df)\nprint(\"Categorical Features:\", categorical)\nprint(\"Non-Categorical Features:\", non_categorical)\nprint(\"Discrete Features:\", discrete)\nprint(\"Continuous Features:\", continuous)\nprint()\nfor i in categorical:\n    print(df[i].value_counts())\n    print()","metadata":{"execution":{"iopub.status.busy":"2025-01-28T02:35:01.130956Z","iopub.execute_input":"2025-01-28T02:35:01.131342Z","iopub.status.idle":"2025-01-28T02:35:01.197809Z","shell.execute_reply.started":"2025-01-28T02:35:01.131309Z","shell.execute_reply":"2025-01-28T02:35:01.196797Z"},"trusted":true,"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"desc=df.describe()\ndesc.loc['unique_count'] = [len(df[col].unique()) for col in desc.columns]\nint_index=['count', 'min', '25%', '50%', '75%', 'max', 'unique_count']\nint_columns=['eeg_id', 'eeg_sub_id', 'eeg_label_offset_seconds', 'spectrogram_id',\n             'spectrogram_sub_id', 'spectrogram_label_offset_seconds', 'label_id',\n             'patient_id', 'seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote',\n             'grda_vote', 'other_vote', 'total_votes']\nfor col in int_columns:\n    desc[col] = desc[col].astype(object)\n    for ix in int_index:\n        desc.at[ix,col] = int(desc.at[ix,col])\ndesc","metadata":{"execution":{"iopub.status.busy":"2025-01-28T02:35:01.19877Z","iopub.execute_input":"2025-01-28T02:35:01.199011Z","iopub.status.idle":"2025-01-28T02:35:01.365362Z","shell.execute_reply.started":"2025-01-28T02:35:01.198991Z","shell.execute_reply":"2025-01-28T02:35:01.36431Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def show_aggregate_distribution(df=df, grp_col='total_votes', agg='count', agg_col='total_votes', return_agg=False):\n    tv = df.groupby(grp_col).agg({agg_col:agg})\n    # fig = plt.figure(figsize=(12,8),)\n    fig = plt.figure()\n    plt.xticks(tv.index)\n    plt.plot(np.array(tv.index), np.array(tv))\n    plt.xlabel('Index')\n    plt.ylabel(agg_col)\n    plt.title(f'{agg_col} by Index')\n    \n    plt.tight_layout()\n    plt.show()\n    if return_agg:\n        return tv\n\nshow_aggregate_distribution()","metadata":{"id":"8ee6e8ce","execution":{"iopub.status.busy":"2025-01-28T02:35:01.366577Z","iopub.execute_input":"2025-01-28T02:35:01.366976Z","iopub.status.idle":"2025-01-28T02:35:01.761441Z","shell.execute_reply.started":"2025-01-28T02:35:01.36694Z","shell.execute_reply":"2025-01-28T02:35:01.760416Z"},"trusted":true,"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Complete the confidence filtering","metadata":{}},{"cell_type":"code","source":"def get_col_ix(col_id): #Consistency checked\n    if isinstance(col_id, int):\n        return 9+col_id\n    elif isinstance(col_id, str):\n        return 9+target_encoding[col_id]\n    \ndef get_vote_col(col_id): #Consistency checked\n    if isinstance(col_id, int):\n        return f'{target_encoding[col_id]}_vote'\n    elif isinstance(col_id, str):\n        return f'{col_id}_vote'\n    \ndef get_normed_ix(col_id): #Consistency checked\n    if isinstance(col_id, int):\n        return 16+col_id\n    elif isinstance(col_id, str):\n        return 16+target_encoding[col_id]\n    \ndef get_normed_col(col_id): #Consistency checked\n    if isinstance(col_id, int):\n        return f'{target_encoding[col_id]}_vote_normed'\n    elif isinstance(col_id, str):\n        return f'{col_id}_vote_normed'","metadata":{"execution":{"iopub.status.busy":"2025-01-28T02:35:01.762439Z","iopub.execute_input":"2025-01-28T02:35:01.762751Z","iopub.status.idle":"2025-01-28T02:35:01.768722Z","shell.execute_reply.started":"2025-01-28T02:35:01.762717Z","shell.execute_reply":"2025-01-28T02:35:01.76787Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df['confidence'] = df.apply(lambda row: row[get_normed_col(target_encoding[row['expert_consensus']])], axis=1)\ndf['full_confidence'] = (df['confidence'] > 0.999) & (df['confidence'] < 1.001)\n\nbins = [i / 10 for i in range(11)]\nbins[-1]+=0.01 \nlabels = [(i+0.5)/10 for i in range(10)]\ndf['confidence_classes'] = pd.cut(df['confidence'], bins=bins, labels=labels, include_lowest=True)\n\ncon_agg = show_aggregate_distribution(df,'confidence_classes', agg_col='confidence_classes', return_agg=True)\n\ntrain_df=df[df['full_confidence']]\ntrain_df.drop(columns=['full_confidence', 'confidence_classes'], inplace=True)\n\ncategorical, non_categorical, discrete, continuous = classify_features(train_df)","metadata":{"execution":{"iopub.status.busy":"2025-01-28T02:35:01.769739Z","iopub.execute_input":"2025-01-28T02:35:01.769991Z","iopub.status.idle":"2025-01-28T02:35:03.040765Z","shell.execute_reply.started":"2025-01-28T02:35:01.769968Z","shell.execute_reply":"2025-01-28T02:35:03.039569Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"tv_agg = show_aggregate_distribution(train_df, return_agg=True)\ntv_agg","metadata":{"execution":{"iopub.status.busy":"2025-01-28T02:35:03.041852Z","iopub.execute_input":"2025-01-28T02:35:03.042262Z","iopub.status.idle":"2025-01-28T02:35:03.370882Z","shell.execute_reply.started":"2025-01-28T02:35:03.042233Z","shell.execute_reply":"2025-01-28T02:35:03.369776Z"},"trusted":true,"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df.to_csv('fc_train_df.csv')\ntrain_df.head()","metadata":{"execution":{"iopub.status.busy":"2025-01-28T02:35:03.372068Z","iopub.execute_input":"2025-01-28T02:35:03.372407Z","iopub.status.idle":"2025-01-28T02:35:03.849412Z","shell.execute_reply.started":"2025-01-28T02:35:03.372381Z","shell.execute_reply":"2025-01-28T02:35:03.848526Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"votes=['seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']\nnormed_votes_columns=[col+\"_normed\" for col in votes]\nvotes_map={k.split('_')[0]:i for i,k in enumerate(votes)}\nvotes_map.update({i:k.split('_')[0] for i,k in enumerate(votes)})\nall_class_dfs=[]\nselect_full_confidence_only=False\nfor vote in votes:\n    if select_full_confidence_only:\n        all_class_dfs.append(train_df[train_df[vote]>0.01])\n    else:\n        all_class_dfs.append(df[df[vote]>0.01])\n\ncounts=[class_df.shape[0] for class_df in all_class_dfs]\nclass_id=[0, 1, 2, 3, 4, 5]\nplt.plot(class_id, counts)","metadata":{"execution":{"iopub.status.busy":"2025-01-28T02:35:03.850273Z","iopub.execute_input":"2025-01-28T02:35:03.850627Z","iopub.status.idle":"2025-01-28T02:35:04.067106Z","shell.execute_reply.started":"2025-01-28T02:35:03.850594Z","shell.execute_reply":"2025-01-28T02:35:04.065961Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"eeg_cols=['Fp1', 'F3', 'C3', 'P3', 'F7', 'T3', 'T5', 'O1', 'Fz', 'Cz', 'Pz', 'Fp2', 'F4', 'C4', 'P4', 'F8', 'T4', 'T6', 'O2', 'EKG']","metadata":{"execution":{"iopub.status.busy":"2025-01-28T02:35:04.068181Z","iopub.execute_input":"2025-01-28T02:35:04.068457Z","iopub.status.idle":"2025-01-28T02:35:04.073297Z","shell.execute_reply.started":"2025-01-28T02:35:04.068433Z","shell.execute_reply":"2025-01-28T02:35:04.072353Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_eeg_path(id):\n    eeg_path=eeg_path = f\"/kaggle/input/hms-harmful-brain-activity-classification/train_eegs/{id}.parquet\"\n    return eeg_path\n\ndef get_eeg(id):\n    eeg_path = f\"/kaggle/input/hms-harmful-brain-activity-classification/train_eegs/{id}.parquet\"\n    eeg = pd.read_parquet(eeg_path)\n    return eeg\n\ndef get_eeg_sample(ix, meta_df, return_normalized_votes=False):\n    id=meta_df.loc[ix, 'eeg_id']\n    # try:\n    #     id=meta_df.loc[ix, 'eeg_id']\n    # except:\n    #     print(ix)\n    eeg=get_eeg(id)\n    start_frame=int(meta_df.loc[ix, 'eeg_label_offset_seconds']*200)\n    end_frame=start_frame+10000\n    eeg_sample=torch.from_numpy(eeg.iloc[start_frame: end_frame].to_numpy())\n    if not return_normalized_votes:\n        return eeg_sample\n    normalized_votes = torch.from_numpy(meta_df.loc[ix, normed_votes_columns].to_numpy(dtype=np.float32))\n    return normalized_votes, eeg_sample\n\ndef download_file(path, download_file_name):\n    os.chdir('./')\n    zip_name = f\"/kaggle/working/{download_file_name}.zip\"\n    command = f\"zip {zip_name} {path} -r\"\n    result = subprocess.run(command, shell=True, capture_output=True, text=True)\n    if result.returncode != 0:\n        print(\"Unable to run zip command!\")\n        print(result.stderr)\n        return \n    display(FileLink(f'{download_file_name}.zip'))\n\n# Define the function to be executed in parallel\ndef eeg_sample_isnan(ix, meta_df):\n    eeg_sample = get_eeg_sample(ix, meta_df).transpose(0, 1)\n    scaler = StandardScaler()\n    eeg_sample = scaler.fit_transform(eeg_sample.T).T\n    if np.isnan(eeg_sample).any():\n        return ix\n    return None\n    \ndef prune_nans(meta_df):\n    # Main code with concurrent futures\n    total_unpruned_rows = meta_df.shape[0]\n    nan_indices = set()\n    \n    # Using ThreadPoolExecutor for concurrent processing\n    with ThreadPoolExecutor(max_workers=4) as executor:\n        # Map the function to the indices, it will run in parallel\n        # print(isinstance(meta_df, pd.DataFrame))\n        results = list(tqdm(executor.map(lambda idx: eeg_sample_isnan(idx, meta_df), range(total_unpruned_rows)),total=total_unpruned_rows))\n    \n    # Collect all the indices where NaNs were found\n    nan_indices = set(filter(None, results))\n    ans_df=meta_df.loc[meta_df.index[~meta_df.index.isin(nan_indices)]].reset_index()\n    return ans_df","metadata":{"execution":{"iopub.status.busy":"2025-01-28T02:35:04.074184Z","iopub.execute_input":"2025-01-28T02:35:04.074471Z","iopub.status.idle":"2025-01-28T02:35:04.09108Z","shell.execute_reply.started":"2025-01-28T02:35:04.074436Z","shell.execute_reply":"2025-01-28T02:35:04.090149Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"uniform_weights=False\nif uniform_weights:\n    min_count=min(counts)\n    print(min_count)\n    frac=0.5\n    len_samples=int(min_count*frac)\n    all_samples=[sample_df.sample(n=len_samples, replace=True) for i,sample_df in enumerate(all_class_dfs)]\n    meta_df=pd.concat(all_samples, ignore_index=True)\n    total_count=sum(counts)\n    weights = [count/total_count for count in counts]\nelif select_full_confidence_only:\n    meta_df=df = train_df.sample(frac=1).reset_index(drop=True)\n    total_count=sum(counts)\n    weights = [count/total_count for count in counts]\nelse:\n    sample_fraction=1\n    meta_df = df.sample(frac=sample_fraction).reset_index(drop=True)\n    total_count=sum(counts)\n    weights = [count/total_count for count in counts]\n\nprint(f\"meta_df.shape = \", meta_df.shape)\nmeta_df=meta_df.sample(frac=0.01, ignore_index=True)\n# meta_df=meta_df.sample(frac=1, ignore_index=True)\nmeta_df=prune_nans(meta_df)\nmeta_df.to_csv('train_df.csv')","metadata":{"execution":{"iopub.status.busy":"2025-01-28T02:35:04.091886Z","iopub.execute_input":"2025-01-28T02:35:04.092164Z","iopub.status.idle":"2025-01-28T02:35:18.459313Z","shell.execute_reply.started":"2025-01-28T02:35:04.092142Z","shell.execute_reply":"2025-01-28T02:35:18.458351Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"num_samples=meta_df.shape[0]\nnum_channels=len(eeg_cols)\nnum_classes=6\nsample_rate=200\nsample_duration=50\nbatch_size=16\nmax_pool_size=1000","metadata":{"execution":{"iopub.status.busy":"2025-01-28T02:35:18.460039Z","iopub.execute_input":"2025-01-28T02:35:18.46031Z","iopub.status.idle":"2025-01-28T02:35:18.4647Z","shell.execute_reply.started":"2025-01-28T02:35:18.460286Z","shell.execute_reply":"2025-01-28T02:35:18.46381Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# modeling_wav2vec2.py variables and imports\nfrom transformers.modeling_outputs import SequenceClassifierOutput\n_HIDDEN_STATES_START_POSITION = 2\n\nclass ModelLoader(Wav2Vec2ForSequenceClassification):\n    loss_fct = None\n\n    @classmethod\n    def from_pretrained(\n        cls,\n        pretrained_model_name_or_path,\n        *model_args,\n        config = None,\n        cache_dir = None,\n        ignore_mismatched_sizes = False,\n        force_download = False,\n        local_files_only = False,\n        token = None,\n        revision = \"main\",\n        use_safetensors = None,\n        # weights_only = True,\n        criterion = None,\n        **kwargs,\n    ):\n        if criterion is None:\n            cls.loss_fct = CrossEntropyLoss()\n        else:\n            cls.loss_fct = criterion\n        return Wav2Vec2ForSequenceClassification.from_pretrained(\n            pretrained_model_name_or_path,\n            *model_args,\n            config = None,\n            cache_dir = None,\n            ignore_mismatched_sizes = False,\n            force_download = False,\n            local_files_only = False,\n            token = None,\n            revision = \"main\",\n            use_safetensors = None,\n            # weights_only = True,\n            **kwargs,\n        )\n\n    def forward(\n        self,\n        input_values,\n        attention_mask = None,\n        output_attentions = None,\n        output_hidden_states = None,\n        return_dict = None,\n        labels = None,\n    ):\n        r\"\"\"\n        labels (`torch.LongTensor` of shape `(batch_size,)`, *optional*):\n            Labels for computing the sequence classification/regression loss. Indices should be in `[0, ...,\n            config.num_labels - 1]`. If `config.num_labels == 1` a regression loss is computed (Mean-Square loss), If\n            `config.num_labels > 1` a classification loss is computed (Cross-Entropy).\n        \"\"\"\n\n        return_dict = return_dict if return_dict is not None else self.config.use_return_dict\n        output_hidden_states = True if self.config.use_weighted_layer_sum else output_hidden_states\n\n        outputs = self.wav2vec2(\n            input_values,\n            attention_mask=attention_mask,\n            output_attentions=output_attentions,\n            output_hidden_states=output_hidden_states,\n            return_dict=return_dict,\n        )\n\n        if self.config.use_weighted_layer_sum:\n            hidden_states = outputs[_HIDDEN_STATES_START_POSITION]\n            hidden_states = torch.stack(hidden_states, dim=1)\n            norm_weights = nn.functional.softmax(self.layer_weights, dim=-1)\n            hidden_states = (hidden_states * norm_weights.view(-1, 1, 1)).sum(dim=1)\n        else:\n            hidden_states = outputs[0]\n\n        hidden_states = self.projector(hidden_states)\n        if attention_mask is None:\n            pooled_output = hidden_states.mean(dim=1)\n        else:\n            padding_mask = self._get_feature_vector_attention_mask(hidden_states.shape[1], attention_mask)\n            expand_padding_mask = padding_mask.unsqueeze(-1).repeat(1, 1, hidden_states.shape[2])\n            hidden_states[~expand_padding_mask] = 0.0\n            pooled_output = hidden_states.sum(dim=1) / padding_mask.sum(dim=1).view(-1, 1)\n\n        logits = self.classifier(pooled_output)\n\n        loss = None\n        if labels is not None:\n            # loss_fct = CrossEntropyLoss()\n            loss = ModelLoader.loss_fct(logits.view(-1, self.config.num_labels), labels.view(-1))\n\n        if not return_dict:\n            output = (logits,) + outputs[_HIDDEN_STATES_START_POSITION:]\n            return ((loss,) + output) if loss is not None else output\n\n        return SequenceClassifierOutput(\n            loss=loss,\n            logits=logits,\n            hidden_states=outputs.hidden_states,\n            attentions=outputs.attentions,\n        )\n                ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-28T02:35:18.465474Z","iopub.execute_input":"2025-01-28T02:35:18.465674Z","iopub.status.idle":"2025-01-28T02:35:18.483404Z","shell.execute_reply.started":"2025-01-28T02:35:18.465655Z","shell.execute_reply":"2025-01-28T02:35:18.482501Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Model and Processor\nprocessor = Wav2Vec2Processor.from_pretrained(\"facebook/wav2vec2-base\")\n# model = Wav2Vec2ForSequenceClassification.from_pretrained(\n#     \"facebook/wav2vec2-base\",\n#     num_labels=num_classes\n# )\nmodel = ModelLoader.from_pretrained(\n    pretrained_model_name_or_path = \"facebook/wav2vec2-base\",\n    num_labels = num_classes,\n    criterion = FocalLoss(gamma=0.7, weights=torch.tensor(weights)),\n)\n\nif os.path.exists(model_folder_path):\n    print(\"Saved model loaded\")\n    model = ModelLoader.from_pretrained(\n    pretrained_model_name_or_path = model_path,\n    num_labels=num_classes\n    )    \n\n# Freeze feature extractor layers\nfor param in model.wav2vec2.feature_extractor.parameters():\n    param.requires_grad = False\n\n# Save the model\nprocessor.save_pretrained('kaggle/working/processor')\ntime.sleep(3)\ndownload_file(\"kaggle/working/processor\",\"processor\")","metadata":{"execution":{"iopub.status.busy":"2025-01-28T02:35:18.487352Z","iopub.execute_input":"2025-01-28T02:35:18.487616Z","iopub.status.idle":"2025-01-28T02:35:25.709589Z","shell.execute_reply.started":"2025-01-28T02:35:18.487589Z","shell.execute_reply":"2025-01-28T02:35:25.708483Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class MetaPoolDataset(Dataset):\n    def __init__(self, meta_df, pool_size, name):\n        self.meta_df = meta_df\n        self.pool_size = pool_size  # Maximum number of samples to load at once\n        self.scaler = StandardScaler()\n        self.name=name\n        self.total_size=0\n        \n        self.current_pool=None\n        self.current_label=None\n        self.current_normalized_votes=None\n        \n        self.index_offset = {}\n        self.start_indices = []\n        self.start_of_file = {}\n        self.end_of_file = {}\n\n    def create_pools(self):\n        max_index=self.meta_df.shape[0]\n        start_index=0\n        file_id=0\n        \n        start_idx = 0\n        meta_start_idx = 0\n        while(meta_start_idx<max_index):\n            # Limit end index to the size of meta_df\n            meta_end_idx = min(meta_start_idx + self.pool_size, len(self.meta_df))\n            \n            file_name=f\"{self.name}_{file_id}\"\n            self.index_offset[file_name]=start_index\n            \n            print(f\"Creating {file_name}\")\n            t1=time.time()\n            pooled_data = np.zeros((meta_end_idx - meta_start_idx, num_channels, 10000))\n            labels = np.zeros((meta_end_idx - meta_start_idx,))\n            normalized_votes = np.zeros((meta_end_idx - meta_start_idx, num_classes))\n#             for i, ix in enumerate(tqdm(self.meta_df.index[meta_start_idx:meta_end_idx])):\n            for i, ix in enumerate(self.meta_df.index[meta_start_idx:meta_end_idx]):\n                eeg_normalized_vote, eeg_sample = get_eeg_sample(ix, self.meta_df, return_normalized_votes=True)\n                eeg_sample = eeg_sample.transpose(0, 1)\n                \n                eeg_sample = self.scaler.fit_transform(eeg_sample.T).T\n                pooled_data[i] = eeg_sample\n                labels[i] = votes_map[self.meta_df.loc[ix, 'expert_consensus'].lower()]\n                normalized_votes[i] = eeg_normalized_vote\n            \n            # Filter out any samples with NaN values\n            mask = ~np.isnan(pooled_data).any(axis=(1, 2))\n            pooled_data = pooled_data[mask]\n            labels = labels[mask]\n            \n            self.start_indices.append(start_index)\n            self.start_of_file[file_name]=meta_start_idx\n            self.end_of_file[file_name]=meta_end_idx\n            \n            self.total_size+=pooled_data.shape[0]\n            start_index+=pooled_data.shape[0]\n            meta_start_idx+=self.pool_size\n            file_id+=1\n            print(f\"Time taken = {(time.time()-t1)/60:.2f} minutes\")\n            print(\"self.start_of_file = \", self.start_of_file)\n            print(\"self.total_size = \", self.total_size)\n            print()\n            \n            del mask, pooled_data, labels\n        \n    def get_file_name(self, idx):\n        assert(self.start_indices is not None)\n        pos=bisect_left(self.start_indices, idx)\n        file_id=pos\n        if pos==len(self.start_indices):\n            file_id-=1\n        elif self.start_indices[pos]>idx:\n            file_id-=1\n        return f\"{self.name}_{file_id}\"\n    \n    def get_current_pool(self, idx):\n        start=end=file_name=pooled_data=labels=eeg_sample=i=ix=None\n        file_name=self.get_file_name(idx)\n        start=self.start_of_file[file_name]\n        end=self.end_of_file[file_name]\n\n        pooled_data = np.zeros((end - start, num_channels, 10000))\n        labels = np.zeros((end - start,))\n        normalized_votes = np.zeros((end - start, num_classes))\n        \n        for i, ix in enumerate(self.meta_df.index[start:end]):\n            eeg_normalized_vote, eeg_sample = get_eeg_sample(ix, self.meta_df, return_normalized_votes=True)\n            eeg_sample = eeg_sample.transpose(0,1)\n            eeg_sample = self.scaler.fit_transform(eeg_sample.T).T\n            pooled_data[i] = eeg_sample\n            labels[i] = votes_map[self.meta_df.loc[ix, 'expert_consensus'].lower()]\n            normalized_votes[i] = eeg_normalized_vote\n\n        # Filter out any samples with NaN values\n        mask = ~np.isnan(pooled_data).any(axis=(1, 2))\n        self.current_pool = pooled_data[mask]\n        self.current_label = labels[mask]\n        self.current_normalized_votes = normalized_votes[mask]","metadata":{"execution":{"iopub.status.busy":"2025-01-28T02:35:25.711643Z","iopub.execute_input":"2025-01-28T02:35:25.711917Z","iopub.status.idle":"2025-01-28T02:35:25.726492Z","shell.execute_reply.started":"2025-01-28T02:35:25.711887Z","shell.execute_reply":"2025-01-28T02:35:25.725475Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Custom Dataset Class\nclass EEGDataset(MetaPoolDataset):\n    def __init__(self, meta_df, pool_size, name, processor=processor, target_sampling_rate=16000, orig_sampling_rate=200):\n        super().__init__(meta_df, pool_size, name)\n        self.processor = processor\n        self.orig_sampling_rate = orig_sampling_rate\n        self.target_sampling_rate = target_sampling_rate\n        self.first_pool_made=False\n        \n        # self.create_pools()\n        # self.current_file = self.get_file_name(0)\n        # self.get_current_pool(0)\n        print(f\"EEGDataset {name} initialized\\n\")\n        \n    def __len__(self):\n        return self.meta_df.shape[0]\n    \n    def __getitem__(self, idx):\n        ix=idx\n        normalized_vote, eeg_sample = get_eeg_sample(ix, self.meta_df, True)\n        eeg_sample = eeg_sample.transpose(0, 1)\n        eeg_sample = self.scaler.fit_transform(eeg_sample.T).T\n        label = votes_map[self.meta_df.loc[ix, 'expert_consensus'].lower()]\n        inputs = self.processor(eeg_sample, sampling_rate=self.target_sampling_rate, return_tensors=\"pt\", padding=True)\n        inputs[\"labels\"] = torch.tensor(label, dtype=torch.long)\n        inputs[\"normalized_votes\"] = normalized_vote\n        return inputs\n        # inputs=None\n        # file_name=None\n        # unchanged_idx=idx\n        # try:\n        #     file_name = self.get_file_name(idx)\n        #     if file_name != self.current_file:\n        #         self.current_file = file_name\n        #         self.get_current_pool(idx)\n\n        #     idx=idx-self.index_offset[self.current_file]\n        #     data = self.current_pool[idx]\n        #     label = int(self.current_label[idx])\n\n        #     # Flatten to 1D\n        #     data = data.reshape(-1)\n\n        #     # Resample to self.target_sampling_rate\n        #     inputs = self.processor(data, sampling_rate=self.target_sampling_rate, return_tensors=\"pt\", padding=True)\n        #     inputs[\"labels\"] = torch.tensor(label, dtype=torch.long)\n        #     inputs[\"normalized_votes\"] = torch.tensor(self.current_normalized_votes[idx])\n        #     return inputs\n        # except Exception as e:\n        #     print(\"Error: \", e)\n        #     print(\"unchanged_idx = \", unchanged_idx)\n        #     print(\"idx = \", idx)\n        #     print(\"file_name = \", file_name)\n        #     if self.current_pool is not None:\n        #         print(\"self.current_pool.shape = \", self.current_pool.shape)\n        #     else:\n        #         print(\"self.current_pool is None\")\n        #     print(self.index_offset)\n        #     traceback.print_exc()\n        #     print()\n\n# Loaders\nif (not os.path.isdir(dataset_folder_path)):\n    len_meta_df=meta_df.shape[0]\n    test_frac=0.15\n    val_frac=0.15\n    train_meta_df=meta_df.iloc[:int((1-test_frac-val_frac)*len_meta_df),:]\n    val_meta_df=meta_df.iloc[int((1-test_frac-val_frac)*len_meta_df):int((1-test_frac)*len_meta_df),:]\n    test_meta_df=meta_df.iloc[int((1-test_frac)*len_meta_df):, :]\n    pool_size = int((max_pool_size//batch_size)*batch_size)\n    train_meta_df.to_csv(\"train_meta_df.csv\")\n    test_meta_df.to_csv(\"test_meta_df.csv\")\n    val_meta_df.to_csv(\"val_meta_df.csv\")\nelse:\n    train_meta_path=dataset_folder_path+\"/train_meta_df.csv\"\n    val_meta_path=dataset_folder_path+\"/val_meta_df.csv\"\n    test_meta_path=dataset_folder_path+\"/test_meta_df.csv\"\n\n    test_meta_df=pd.read_csv(test_meta_path)\n    train_meta_df=pd.read_csv(train_meta_path)\n    val_meta_df=pd.read_csv(val_meta_path)\n","metadata":{"execution":{"iopub.status.busy":"2025-01-28T02:35:25.727773Z","iopub.execute_input":"2025-01-28T02:35:25.728147Z","iopub.status.idle":"2025-01-28T02:35:26.197603Z","shell.execute_reply.started":"2025-01-28T02:35:25.728114Z","shell.execute_reply":"2025-01-28T02:35:26.196754Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pool_size=(max_pool_size//batch_size)*batch_size\n\ntrain_meta_df.reset_index(drop=True, inplace=True)\nval_meta_df.reset_index(drop=True, inplace=True)\ntest_meta_df.reset_index(drop=True, inplace=True)\n\nif (not os.path.isdir(dataset_folder_path)):\n    train_dataset = EEGDataset(train_meta_df, pool_size, \"train\")\n    val_dataset = EEGDataset(val_meta_df, pool_size, \"val\")\n    test_dataset = EEGDataset(test_meta_df, pool_size, \"test\")\n    all_datasets={\"train_dataset\":train_dataset, \"test_dataset\":test_dataset, \"val_dataset\":val_dataset,}\n    with open('all_datasets.pkl', 'wb') as f:\n        pickle.dump(all_datasets, f)\n    \nelse:\n    train_dataset=EEGDataset(train_meta_df, pool_size, \"train\")\n    test_dataset=EEGDataset(test_meta_df, pool_size, \"test\")\n    val_dataset=EEGDataset(val_meta_df, pool_size, \"val\")\n\ntrain_loader = DataLoader(train_dataset, batch_size=batch_size, prefetch_factor=1, pin_memory=True, num_workers=1)\nval_loader = DataLoader(val_dataset, batch_size=batch_size, prefetch_factor=1, pin_memory=True, num_workers=1)\ntest_loader = DataLoader(test_dataset, batch_size=batch_size, prefetch_factor=1, pin_memory=True, num_workers=1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-28T02:35:26.198561Z","iopub.execute_input":"2025-01-28T02:35:26.198891Z","iopub.status.idle":"2025-01-28T02:35:26.208405Z","shell.execute_reply.started":"2025-01-28T02:35:26.198865Z","shell.execute_reply":"2025-01-28T02:35:26.207536Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"download_file(\"./train_meta_df.csv\", \"train_meta_df\")\ndownload_file(\"./test_meta_df.csv\", \"test_meta_df\")\ndownload_file(\"./val_meta_df.csv\", \"val_meta_df\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-28T02:35:26.209197Z","iopub.execute_input":"2025-01-28T02:35:26.209419Z","iopub.status.idle":"2025-01-28T02:35:26.239948Z","shell.execute_reply.started":"2025-01-28T02:35:26.209399Z","shell.execute_reply":"2025-01-28T02:35:26.239128Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Optimizer and Loss\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel.to(device)\n\noptimizer = torch.optim.AdamW(model.parameters(), lr=1e-5)\n# criterion = nn.CrossEntropyLoss()\n\nnum_epochs=4","metadata":{"execution":{"iopub.status.busy":"2025-01-28T02:35:26.240904Z","iopub.execute_input":"2025-01-28T02:35:26.241219Z","iopub.status.idle":"2025-01-28T02:35:29.49432Z","shell.execute_reply.started":"2025-01-28T02:35:26.241178Z","shell.execute_reply":"2025-01-28T02:35:29.493486Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Done above","metadata":{"execution":{"iopub.status.busy":"2025-01-28T02:35:29.495148Z","iopub.execute_input":"2025-01-28T02:35:29.49546Z","iopub.status.idle":"2025-01-28T02:35:29.499771Z","shell.execute_reply.started":"2025-01-28T02:35:29.495429Z","shell.execute_reply":"2025-01-28T02:35:29.498666Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def make_predictions(data_loader, name, data_preds=None, data_labels=None):\n    if (data_preds is None) and (data_labels is None):\n        data_preds, data_labels = [], []\n        with torch.no_grad():\n            for batch in tqdm(data_loader, desc=f\"Checking model on {name} dataset\"):\n                inputs={}\n                inputs['input_values'] = batch['input_values'].to(torch.float32).to(device)\n                if len(inputs['input_values'].shape)==1:\n                    inputs['input_values'] = torch.unsqueeze(inputs['input_values'], 0)\n                    inputs['input_values'].to(device)\n                inputs['input_values'] = torch.reshape(inputs['input_values'],(-1, 200000))\n                \n                inputs['labels'] = batch['labels'].to(torch.int64).to(device)\n                outputs = model(**inputs)\n\n                logits = outputs.logits\n                preds = torch.argmax(logits, dim=-1)\n                labels = inputs['labels']\n                \n                data_preds.extend(preds.tolist())\n                data_labels.extend(labels.tolist())\n\n    # Calculate metrics\n    data_accuracy = accuracy_score(data_labels, data_preds)\n    data_f1=f1_score(data_labels, data_preds, average='weighted')\n    data_precision = precision_score(data_labels, data_preds, average='weighted')\n    data_recall = recall_score(data_labels, data_preds, average='weighted')\n    print(f\"{name} Dataset - Accuracy: {data_accuracy:.4f} | Precision: {data_precision:.4f} | Recall: {data_recall:.4f} | F1: {data_f1:.4f}\")\n","metadata":{"execution":{"iopub.status.busy":"2025-01-28T02:35:29.500773Z","iopub.execute_input":"2025-01-28T02:35:29.501038Z","iopub.status.idle":"2025-01-28T02:35:29.518073Z","shell.execute_reply.started":"2025-01-28T02:35:29.501015Z","shell.execute_reply":"2025-01-28T02:35:29.517105Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# epoch_offset=int(model_path[-5])\nepoch_offset=2","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-28T02:35:29.518972Z","iopub.execute_input":"2025-01-28T02:35:29.519323Z","iopub.status.idle":"2025-01-28T02:35:29.535554Z","shell.execute_reply.started":"2025-01-28T02:35:29.519293Z","shell.execute_reply":"2025-01-28T02:35:29.534534Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class EvaluatePartition:\n    def __init__(self, model, data_loader, name, device=device):\n        self.model=model\n        self.data_loader=data_loader\n        self.name=name\n        self.device=device\n\n        self.calc_softmax=torch.nn.Softmax(dim=-1)  \n        self.data_preds = []\n        self.data_labels = []\n        self.data_probs = []\n        self.total_loss = 0.0\n        self.total_kl_loss = 0.0\n        self.num_classes = None\n        self.final_metrics = {}\n        \n\n    def _reset_base_metrics(self, base_metrics=None):\n        self.data_preds = []\n        self.data_labels = []\n        self.data_probs = []\n        self.total_loss = 0.0\n        self.total_kl_loss = 0.0\n        if base_metrics is not None:\n            self.data_preds = base_metrics[\"data_preds\"]\n            self.data_labels = base_metrics[\"data_labels\"]\n            self.data_probs = base_metrics[\"data_probs\"]\n            self.total_loss = base_metrics[\"total_loss\"]\n            self.total_kl_loss = base_metrics[\"total_kl_loss\"]\n        self.final_metrics.clear()\n    \n    def _evaluate(self, base_metrics=None):\n        self._reset_base_metrics(base_metrics)\n        \n        if base_metrics is None:\n            with torch.no_grad():\n                for batch in tqdm(self.data_loader, desc=f\"Checking model on {self.name} dataset\", ncols=100, leave=True, dynamic_ncols=True):\n                    inputs=self._get_input_kwargs(batch)\n                    normalized_votes=inputs['normalized_votes']\n                    del inputs['normalized_votes']\n                    \n                    outputs = self.model(**inputs)\n                    \n                    inputs['normalized_votes']=normalized_votes\n                    loss=outputs.loss\n                    self.total_loss+=loss.item()\n                    logits, preds, probs, labels = self._process_model_outputs(inputs, outputs)\n                \n        return self._get_final_metrics()\n    \n    def _get_input_kwargs(self, batch):\n        inputs = {}\n        inputs['input_values'] = batch['input_values'].to(torch.float32).to(self.device)\n        if len(inputs['input_values'].shape) == 1:\n            inputs['input_values'] = torch.unsqueeze(inputs['input_values'], 0)\n            inputs['input_values'].to(self.device)\n        inputs['input_values'] = torch.reshape(inputs['input_values'], (-1, num_channels*10000))\n        inputs['labels'] = batch['labels'].to(torch.int64).to(self.device)\n        inputs['normalized_votes']=batch['normalized_votes'].to(self.device)\n        return inputs\n\n    def _process_model_outputs(self, inputs, outputs):\n        logits = outputs.logits \n        # Collect predictions and labels\n        preds = torch.argmax(logits, dim=-1)\n        self.data_preds.extend(preds.tolist())\n\n        labels = inputs['labels']\n        self.data_labels.extend(labels.tolist())\n        \n        probs = self.calc_softmax(logits)\n        self.data_probs.append(probs.detach().cpu().numpy())\n\n        # print(\"probabilities = \", F.log_softmax(logits, dim=-1))\n        # print(\"log probabilities = \", torch.log(F.log_softmax(logits, dim=-1)))\n        kl_loss = F.kl_div(F.log_softmax(logits, dim=-1), inputs['normalized_votes'], reduction='batchmean')\n        self.total_kl_loss += kl_loss.item()\n        \n        return logits, preds, probs, labels\n\n    def _get_final_metrics(self, base_metrics=None):\n        self._finalize_epoch_data()\n        self.num_classes = self.data_probs.shape[1]\n        self._add_regressive_scores()\n        self._add_evaluation_graphs()\n        return self.final_metrics\n        \n    def _finalize_epoch_data(self):\n        self.data_labels = torch.tensor(self.data_labels).cpu().numpy()\n        self.data_preds = torch.tensor(self.data_preds).cpu().numpy()\n        if isinstance(self.data_probs, list):\n            self.data_probs = np.concatenate(self.data_probs, axis=0)\n        self.data_probs = torch.tensor(self.data_probs).cpu().numpy()\n        \n    def _add_regressive_scores(self):\n        data_accuracy = accuracy_score(self.data_labels, self.data_preds)\n        data_f1 = f1_score(self.data_labels, self.data_preds, average='weighted')\n        data_precision = precision_score(self.data_labels, self.data_preds, average='weighted')\n        data_recall = recall_score(self.data_labels, self.data_preds, average='weighted')\n        cohen_kappa = cohen_kappa_score(self.data_labels, self.data_preds)\n        \n        regressive_scores={\n            \"accuracy\": data_accuracy,\n            \"f1_score\": data_f1,\n            \"precision\": data_precision,\n            \"recall\": data_recall,\n            \"cohen_kappa\": cohen_kappa,\n            \"total_loss\": self.total_loss,\n            \"total_kl_loss\": self.total_kl_loss,\n        }\n        print(f\"{self.name} Dataset - Accuracy: {data_accuracy:.4f} | Precision: {data_precision:.4f} | Recall: {data_recall:.4f} | F1: {data_f1:.4f}\")\n        print(f\"Total Loss: {self.total_loss:.4f} | Total KL Loss: {self.total_kl_loss:.4f} | cohen_kappa: {cohen_kappa:.4f}\")\n        self.final_metrics.update(regressive_scores)\n        \n    def _add_evaluation_graphs(self):\n        try:\n            auc_roc = roc_auc_score(self.data_labels, self.data_probs, multi_class='ovr', average='weighted')\n        except ValueError:\n            auc_roc = None\n    \n        pr_curves = {}\n        for i in range(self.num_classes):\n            precision, recall, _ = precision_recall_curve(self.data_labels == i, self.data_probs[:, i])\n            pr_curves[f\"class_{i}\"] = {\"precision\": precision.tolist(), \"recall\": recall.tolist()}\n        class_report = classification_report(self.data_labels, self.data_preds, output_dict=True)\n        \n        evaluation_curves={\n            \"auc_roc\": auc_roc,\n            \"support\": class_report[\"weighted avg\"][\"support\"],\n            \"pr_curves\": pr_curves\n        }\n        self.final_metrics.update(evaluation_curves)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-28T02:35:29.5366Z","iopub.execute_input":"2025-01-28T02:35:29.536981Z","iopub.status.idle":"2025-01-28T02:35:29.986409Z","shell.execute_reply.started":"2025-01-28T02:35:29.53693Z","shell.execute_reply":"2025-01-28T02:35:29.985043Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # Training Loop - EvaluatePartition based\nexperiment_data={}\n\ntraining = EvaluatePartition(model = model, data_loader = train_loader, name=\"train\", device=device)\nvalidation = EvaluatePartition(model = model, data_loader = val_loader, name=\"val\", device=device)\ntesting = EvaluatePartition(model = model, data_loader = test_loader, name=\"test\", device=device)\n\nfor epoch_id in range(num_epochs):\n    epoch=epoch_id+epoch_offset\n    model.train()\n    total_loss = 0.0\n    total_kl_loss = 0.0\n    train_preds, train_labels, train_probs = [], [], []\n\n    training._reset_base_metrics()\n    print(f\"For Epoch: {epoch}\")\n    epoch_data={}\n    for i, batch in enumerate(tqdm(train_loader)):\n        inputs = training._get_input_kwargs(batch)\n        normalized_votes=inputs['normalized_votes']\n        del inputs['normalized_votes']\n        \n        outputs = model(**inputs)\n        inputs['normalized_votes']=normalized_votes\n        \n        loss = outputs.loss\n        training.total_loss += loss.item()\n        loss.backward()\n        optimizer.step()\n        optimizer.zero_grad()        \n\n        training._process_model_outputs(inputs, outputs)\n        \n\n    train_metrics = training._get_final_metrics()\n    val_metrics = validation._evaluate()\n    test_metrics = testing._evaluate()\n\n    epoch_data.update(train_metrics)\n    epoch_data.update(val_metrics)\n    epoch_data.update(test_metrics)\n    experiment_data[f\"epoch_{epoch}\"] = {}\n    experiment_data[f\"epoch_{epoch}\"].update(epoch_data)\n    \n    model.save_pretrained(f\"model_{epoch}.pkl\")\n    download_file(f\"./model_{epoch}.pkl\", f\"model_{epoch}.pkl\")\n    print(\"Model saved\")\n    \n    print()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-28T02:35:29.987594Z","iopub.execute_input":"2025-01-28T02:35:29.987987Z","iopub.status.idle":"2025-01-28T11:15:00.788959Z","shell.execute_reply.started":"2025-01-28T02:35:29.987948Z","shell.execute_reply":"2025-01-28T11:15:00.78781Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"final_data_path = \"./experiment_data.pkl\"\nwith open(final_data_path, 'wb') as f:\n    pickle.dump(experiment_data, f)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-28T11:15:00.7904Z","iopub.execute_input":"2025-01-28T11:15:00.790814Z","iopub.status.idle":"2025-01-28T11:15:00.803138Z","shell.execute_reply.started":"2025-01-28T11:15:00.790772Z","shell.execute_reply":"2025-01-28T11:15:00.802353Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # Training Loop - default\n# for epoch in range(num_epochs):\n#     model.train()\n#     total_loss = 0.0\n#     train_preds, train_labels = [], []\n#     print(f\"For Epoch: {epoch}\")\n#     debug=False\n#     for i, batch in enumerate(tqdm(train_loader)):\n#         inputs={}\n#         inputs['input_values'] = batch['input_values'].to(torch.float32).to(device)\n#         if len(inputs['input_values'].shape)==1:\n#             inputs['input_values'] = torch.unsqueeze(inputs['input_values'], 0)\n#             inputs['input_values'].to(device)\n#         inputs['input_values'] = torch.reshape(inputs['input_values'],(-1, 200000))\n        \n#         inputs['labels'] = batch['labels'].to(torch.int64).to(device)\n#         # print(inputs['input_values'].device)\n#         # print(inputs['input_values'].shape)\n        \n#         if debug:\n#             break\n#         outputs = model(**inputs)\n#         logits = outputs.logits\n#         preds = torch.argmax(logits, dim=-1)\n#         labels=batch[\"labels\"].to(device)\n#         train_preds.extend(preds.tolist())\n#         train_labels.extend(labels.tolist())\n        \n#         loss = outputs.loss\n#         total_loss += loss.item()\n#         loss.backward()\n#         optimizer.step()\n#         optimizer.zero_grad()\n    \n#     if debug:\n#         break\n#     make_predictions(train_loader, \"train\", train_preds, train_labels)\n#     make_predictions(val_loader, \"val\")\n#     make_predictions(test_loader, \"test\")\n#     model.save_pretrained(f\"model_{epoch+epoch_offset}.pkl\")\n#     download_file(f\"./model_{epoch+epoch_offset}.pkl\", f\"model_{epoch+epoch_offset}.pkl\")\n#     print(\"Model saved\")\n    \n#     print()\n","metadata":{"execution":{"iopub.status.busy":"2025-01-28T11:15:00.804179Z","iopub.execute_input":"2025-01-28T11:15:00.804493Z","iopub.status.idle":"2025-01-28T11:15:00.809148Z","shell.execute_reply.started":"2025-01-28T11:15:00.804461Z","shell.execute_reply":"2025-01-28T11:15:00.808193Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}