Chantland/HRAF_Multilabel_SubClasses
011
1{2 "cells": [3 {4 "cell_type": "code",5 "execution_count": 75,6 "metadata": {},7 "outputs": [],8 "source": [9 "from datasets.dataset_dict import DatasetDict\n",10 "from datasets import Dataset, concatenate_datasets\n",11 "import evaluate\n",12 "import os\n",13 "import json\n",14 "import pandas as pd\n",15 "import numpy as np\n",16 "\n",17 "from sklearn.model_selection import train_test_split \n"18 ]19 },20 {21 "cell_type": "code",22 "execution_count": 76,23 "metadata": {},24 "outputs": [],25 "source": [26 "# from huggingface_hub import notebook_login\n",27 "# # # copy and paste this code in the terminal: huggingface-cli login \n",28 "# # # then paste this token: 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# It may not show up but still paste the toke in and press enter\n",29 "\n",30 "\n",31 "# notebook_login()"32 ]33 },34 {35 "attachments": {},36 "cell_type": "markdown",37 "metadata": {},38 "source": [39 "## Import the dataset"40 ]41 },42 {43 "cell_type": "code",44 "execution_count": 89,45 "metadata": {},46 "outputs": [47 {48 "data": {49 "text/html": [50 "<div>\n",51 "<style scoped>\n",52 " .dataframe tbody tr th:only-of-type {\n",53 " vertical-align: middle;\n",54 " }\n",55 "\n",56 " .dataframe tbody tr th {\n",57 " vertical-align: top;\n",58 " }\n",59 "\n",60 " .dataframe thead tr th {\n",61 " text-align: left;\n",62 " }\n",63 "</style>\n",64 "<table border=\"1\" class=\"dataframe\">\n",65 " <thead>\n",66 " <tr>\n",67 " <th></th>\n",68 " <th colspan=\"10\" halign=\"left\">CULTURE</th>\n",69 " <th>...</th>\n",70 " <th colspan=\"3\" halign=\"left\">ACTION</th>\n",71 " <th>OTHER</th>\n",72 " <th colspan=\"3\" halign=\"left\">CODER</th>\n",73 " <th>OTHER</th>\n",74 " <th colspan=\"2\" halign=\"left\">CODER</th>\n",75 " </tr>\n",76 " <tr>\n",77 " <th></th>\n",78 " <th>Passage Number</th>\n",79 " <th>Region</th>\n",80 " <th>SubRegion</th>\n",81 " <th>Culture</th>\n",82 " <th>DocTitle</th>\n",83 " <th>Section</th>\n",84 " <th>Author</th>\n",85 " <th>Page</th>\n",86 " <th>Year</th>\n",87 " <th>OCM</th>\n",88 " <th>...</th>\n",89 " <th>Other</th>\n",90 " <th>Description</th>\n",91 " <th>Local_terms</th>\n",92 " <th>Other_Comments</th>\n",93 " <th>Run_Number</th>\n",94 " <th>Finished</th>\n",95 " <th>Coder</th>\n",96 " <th>Other_Comments.1</th>\n",97 " <th>Dataset</th>\n",98 " <th>Info</th>\n",99 " </tr>\n",100 " </thead>\n",101 " <tbody>\n",102 " <tr>\n",103 " <th>0</th>\n",104 " <td>1392</td>\n",105 " <td>Asia</td>\n",106 " <td>South Asia</td>\n",107 " <td>Andamans</td>\n",108 " <td>Hygiene and medical practices among the Onge (...</td>\n",109 " <td>1. Habitation</td>\n",110 " <td>Cipriani, Lidio</td>\n",111 " <td>484</td>\n",112 " <td>1961</td>\n",113 " <td>['171', '301', '727', '751', '765', '775', '777']</td>\n",114 " <td>...</td>\n",115 " <td>1</td>\n",116 " <td>Several customs are believed to connect with t...</td>\n",117 " <td>ibidanghe: made from decorated human jawbone ...</td>\n",118 " <td>General note of this spreadsheet - many of the...</td>\n",119 " <td>1</td>\n",120 " <td>True</td>\n",121 " <td>YM</td>\n",122 " <td>NaN</td>\n",123 " <td>1</td>\n",124 " <td>Dataset 1: ['750', '751', '752', '753'] Coun...</td>\n",125 " </tr>\n",126 " <tr>\n",127 " <th>1</th>\n",128 " <td>1393</td>\n",129 " <td>Asia</td>\n",130 " <td>South Asia</td>\n",131 " <td>Andamans</td>\n",132 " <td>Hygiene and medical practices among the Onge (...</td>\n",133 " <td>3. Food</td>\n",134 " <td>Cipriani, Lidio</td>\n",135 " <td>487</td>\n",136 " <td>1961</td>\n",137 " <td>['136', '231', '271', '312', '415', '516', '751']</td>\n",138 " <td>...</td>\n",139 " <td>0</td>\n",140 " <td>No action is mentioned</td>\n",141 " <td>0</td>\n",142 " <td>NaN</td>\n",143 " <td>1</td>\n",144 " <td>True</td>\n",145 " <td>YM</td>\n",146 " <td>NaN</td>\n",147 " <td>1</td>\n",148 " <td>Dataset 2: ['784', '731', '732', '777', '791',...</td>\n",149 " </tr>\n",150 " <tr>\n",151 " <th>2</th>\n",152 " <td>1395</td>\n",153 " <td>Asia</td>\n",154 " <td>South Asia</td>\n",155 " <td>Andamans</td>\n",156 " <td>Hygiene and medical practices among the Onge (...</td>\n",157 " <td>3. Food</td>\n",158 " <td>Cipriani, Lidio</td>\n",159 " <td>490</td>\n",160 " <td>1961</td>\n",161 " <td>['114', '137', '164', '262', '273', '751', '825']</td>\n",162 " <td>...</td>\n",163 " <td>0</td>\n",164 " <td>Certain foods, such as Pteropus (giant bat) an...</td>\n",165 " <td>Pteropus: a giant bat eaten by the Andaman Is...</td>\n",166 " <td>NaN</td>\n",167 " <td>1</td>\n",168 " <td>True</td>\n",169 " <td>YM</td>\n",170 " <td>NaN</td>\n",171 " <td>1</td>\n",172 " <td>Run 1: Spring 2023 Coding of Sickness dataset ...</td>\n",173 " </tr>\n",174 " </tbody>\n",175 "</table>\n",176 "<p>3 rows × 43 columns</p>\n",177 "</div>"178 ],179 "text/plain": [180 " CULTURE \\\n",181 " Passage Number Region SubRegion Culture \n",182 "0 1392 Asia South Asia Andamans \n",183 "1 1393 Asia South Asia Andamans \n",184 "2 1395 Asia South Asia Andamans \n",185 "\n",186 " \\\n",187 " DocTitle Section \n",188 "0 Hygiene and medical practices among the Onge (... 1. Habitation \n",189 "1 Hygiene and medical practices among the Onge (... 3. Food \n",190 "2 Hygiene and medical practices among the Onge (... 3. Food \n",191 "\n",192 " \\\n",193 " Author Page Year \n",194 "0 Cipriani, Lidio 484 1961 \n",195 "1 Cipriani, Lidio 487 1961 \n",196 "2 Cipriani, Lidio 490 1961 \n",197 "\n",198 " ... ACTION \\\n",199 " OCM ... Other \n",200 "0 ['171', '301', '727', '751', '765', '775', '777'] ... 1 \n",201 "1 ['136', '231', '271', '312', '415', '516', '751'] ... 0 \n",202 "2 ['114', '137', '164', '262', '273', '751', '825'] ... 0 \n",203 "\n",204 " \\\n",205 " Description \n",206 "0 Several customs are believed to connect with t... \n",207 "1 No action is mentioned \n",208 "2 Certain foods, such as Pteropus (giant bat) an... \n",209 "\n",210 " \\\n",211 " Local_terms \n",212 "0 ibidanghe: made from decorated human jawbone ... \n",213 "1 0 \n",214 "2 Pteropus: a giant bat eaten by the Andaman Is... \n",215 "\n",216 " OTHER CODER \\\n",217 " Other_Comments Run_Number Finished \n",218 "0 General note of this spreadsheet - many of the... 1 True \n",219 "1 NaN 1 True \n",220 "2 NaN 1 True \n",221 "\n",222 " OTHER CODER \\\n",223 " Coder Other_Comments.1 Dataset \n",224 "0 YM NaN 1 \n",225 "1 YM NaN 1 \n",226 "2 YM NaN 1 \n",227 "\n",228 " \n",229 " Info \n",230 "0 Dataset 1: ['750', '751', '752', '753'] Coun... \n",231 "1 Dataset 2: ['784', '731', '732', '777', '791',... \n",232 "2 Run 1: Spring 2023 Coding of Sickness dataset ... \n",233 "\n",234 "[3 rows x 43 columns]"235 ]236 },237 "execution_count": 89,238 "metadata": {},239 "output_type": "execute_result"240 }241 ],242 "source": [243 "\n",244 "df_path = \"../../../eHRAF_Scraper-Analysis-and-Prep/Data/\"\n",245 "dataFolder = r\"(subjects-(contracts_OR_disabilities_OR_disasters_OR_friendships_OR_gift_giving_OR_infant_feeding_OR_lineages_OR_local_officials_OR_luck_and_chance_OR_magicians_and_diviners_OR_mortuary_specialists_OR_nuclear_family_OR_priesthood_OR_prophet/\"\n",246 "# dataFolder = r'subjects-(sickness)_FILTERS-culture_level_samples(PSF)'\n",247 "\n",248 "# Get model and centralized path if relevent\n",249 "model_name = \"HRAF_MultiLabel_SubClasses_Kfolds\"\n",250 "path = f\"\" #Path to centralized file locations (leave blank if centralized location is here)\n",251 "\n",252 "\n",253 "\n",254 "#load df (only load one of these commented out lines)\n",255 "# df = pd.read_excel(f\"{df_path}{dataFolder}/_Altogether_Dataset_RACoded.xlsx\", header=[0,1], index_col=0) # Fall 2023 sickness + non-sickness\n",256 "df = pd.read_excel(f\"{df_path}{dataFolder}/_Altogether_Dataset_RACoded_Combined.xlsx\", header=[0,1], index_col=0) # Spring 2023 - Spring 2024 sickness + nonsickness dataset\n",257 "df.head(3)"258 ]259 },260 {261 "cell_type": "markdown",262 "metadata": {},263 "source": [264 "## Preprocess"265 ]266 },267 {268 "cell_type": "markdown",269 "metadata": {},270 "source": [271 "### Remove Duplicates"272 ]273 },274 {275 "cell_type": "markdown",276 "metadata": {},277 "source": [278 "There were multiple iterations of Research Assistants labeling the data. <br>\n",279 "We will just use run number 1 and 3 with preference to 3 when there are duplicates (as it is the most recent and robust)\n",280 "\n"281 ]282 },283 {284 "cell_type": "code",285 "execution_count": 90,286 "metadata": {},287 "outputs": [288 {289 "data": {290 "text/plain": [291 "Run_Number Dataset\n",292 "1 1 1926\n",293 "2 1 51\n",294 "3 1 4844\n",295 " 2 4184\n",296 "Name: count, dtype: int64"297 ]298 },299 "execution_count": 90,300 "metadata": {},301 "output_type": "execute_result"302 }303 ],304 "source": [305 "### DELETE\n",306 "# Show run number and dataset\n",307 "df[\"CODER\"][[\"Run_Number\", \"Dataset\"]].value_counts(sort=False, dropna=False)"308 ]309 },310 {311 "cell_type": "code",312 "execution_count": 91,313 "metadata": {},314 "outputs": [315 {316 "name": "stdout",317 "output_type": "stream",318 "text": [319 "Duplicate Passages: 0\n"320 ]321 },322 {323 "data": {324 "text/plain": [325 "(CODER, Run_Number)\n",326 "3 9027\n",327 "1 1340\n",328 "Name: count, dtype: int64"329 ]330 },331 "execution_count": 91,332 "metadata": {},333 "output_type": "execute_result"334 }335 ],336 "source": [337 "# useRuns = [1,3] #Only include these runs (NOTE THIS IS COMMENTED OUT IN ORDER TO NOT RUIN THE SUBLABEL DATASET BUT EVENTUALLY YOU SHOULD USE THIS CODE)\n",338 "# df = df.loc[df[(\"CODER\",\"Run_Number\")].isin(useRuns)]\n",339 "\n",340 "mask_NotDuplicate = ~(df.duplicated((\"CULTURE\",\"Passage\"), keep=False)) \n",341 "mask_Dataset2 = df[(\"CODER\",\"Run_Number\")]==3\n",342 "\n",343 "df = df[(mask_NotDuplicate) | (mask_Dataset2)]\n",344 "\n",345 "# Remove certain passages which should not be in training or inference (these are duplicates that had to be manually found by a human)\n",346 "values_to_remove = [3252, 33681, 6758, 10104]\n",347 "df = df[~df[('CULTURE','Passage Number')].isin(values_to_remove)]\n",348 "\n",349 "print(\"Duplicate Passages:\",sum((df.duplicated((\"CULTURE\",\"Passage\")))))\n",350 "df[(\"CODER\",\"Run_Number\")].value_counts()"351 ]352 },353 {354 "cell_type": "markdown",355 "metadata": {},356 "source": [357 "### Set up Dataset\n"358 ]359 },360 {361 "cell_type": "code",362 "execution_count": 92,363 "metadata": {},364 "outputs": [365 {366 "name": "stdout",367 "output_type": "stream",368 "text": [369 "Columns excluded:\n",370 " {('CODER', 'Info'), ('CODER', 'Dataset'), ('CAUSE', 'Local_Terms'), ('CULTURE', 'OWC'), ('CULTURE', 'Region'), ('ACTION', 'No_Info'), ('CULTURE', 'OCM'), ('CULTURE', 'DocTitle'), ('CAUSE', 'Description'), ('CAUSE', 'Just_Happens'), ('CULTURE', 'Author'), ('ACTION', 'Local_terms'), ('CULTURE', 'SubRegion'), ('CULTURE', 'Page'), ('CODER', 'Coder'), ('OTHER', 'Other_Comments.1'), ('EVENT', 'Local_Terms'), ('CULTURE', 'Year'), ('CAUSE', 'No_Info'), ('ACTION', 'Other'), ('EVENT', 'Description'), ('CULTURE', 'Culture'), ('EVENT', 'No_Info'), ('CULTURE', 'Section'), ('CAUSE', 'Other'), ('CODER', 'Finished'), ('CODER', 'Run_Number'), ('ACTION', 'Description'), ('OTHER', 'Other_Comments')} \n",371 "\n",372 "Columns included:\n",373 "ID\n",374 "passage\n",375 "EVENT_Illness\n",376 "EVENT_Accident\n",377 "EVENT_Other\n",378 "CAUSE_Material_Physical\n",379 "CAUSE_Spirits_Gods\n",380 "CAUSE_Witchcraft_Sorcery\n",381 "CAUSE_Rule_Violation_Taboo\n",382 "ACTION_Physical_Material\n",383 "ACTION_Technical_Specialist\n",384 "ACTION_Divination\n",385 "ACTION_Shaman_Medium_Healer\n",386 "ACTION_Priest_High_Religion\n"387 ]388 }389 ],390 "source": [391 "#Construct col list\n",392 "cols = list(df.columns)\n",393 "id_index = cols.index(('CULTURE', \"Passage Number\"))\n",394 "passage_index = cols.index(('CULTURE', \"Passage\"))\n",395 "event_index = cols.index(('EVENT', \"No_Info\"))\n",396 "cause_index = cols.index(('CAUSE', \"No_Info\"))\n",397 "action_index = cols.index(('ACTION', \"No_Info\"))\n",398 "# get a list of all the multi-indexed column names we want to evaluate\n",399 "# col_list = [cols[id_index]] + [cols[passage_index]] + cols[event_index:event_index+4] + cols[cause_index:cause_index+7] + cols[action_index:action_index+7] #to include all columns including No_info\n",400 "col_list = [cols[id_index]] + [cols[passage_index]] + cols[event_index+1:event_index+4] + cols[cause_index+1:cause_index+7] + cols[action_index+1:action_index+7] # to include al columns BUT No_info\n",401 "\n",402 "\n",403 "## Remove the following columns from the dataset. Based on the results of previous models ran and the bias of the categories, remove the following columns from the dataset\n",404 "remv_cols = [(\"CAUSE\",\"Just_Happens\"),(\"CAUSE\",\"Other\"),(\"ACTION\",\"Other\")]\n",405 "for remv in remv_cols:\n",406 " col_list.remove(remv)\n",407 "\n",408 "\n",409 "\n",410 "# get column names to ascribe to the new data frame\n",411 "colNames = [\"ID\",\"passage\"]\n",412 "for category, sub_cat in col_list:\n",413 " # skip passage and id which have already been added\n",414 " # print(category, sub_cat)\n",415 " if category == \"CULTURE\":\n",416 " continue\n",417 " if sub_cat == \"No_Info\":\n",418 " colNames += [category]#this to include main classes, we will hold off on that\n",419 " pass\n",420 " else:\n",421 " colNames += [f'{category}_{sub_cat}']\n",422 "\n",423 "print(\"Columns excluded:\\n\", set(cols)-set(col_list),\"\\n\")\n",424 "# for col in col_list\n",425 "print(\"Columns included:\")\n",426 "for col in colNames:\n",427 " print(col)\n",428 "# colNames\n"429 ]430 },431 {432 "cell_type": "markdown",433 "metadata": {},434 "source": [435 "### Create Huggingface Dataset and do splits"436 ]437 },438 {439 "cell_type": "code",440 "execution_count": 93,441 "metadata": {},442 "outputs": [443 {444 "data": {445 "text/plain": [446 "DatasetDict({\n",447 " train: Dataset({\n",448 " features: ['ID', 'passage', 'EVENT_Illness', 'EVENT_Accident', 'EVENT_Other', 'CAUSE_Material_Physical', 'CAUSE_Spirits_Gods', 'CAUSE_Witchcraft_Sorcery', 'CAUSE_Rule_Violation_Taboo', 'ACTION_Physical_Material', 'ACTION_Technical_Specialist', 'ACTION_Divination', 'ACTION_Shaman_Medium_Healer', 'ACTION_Priest_High_Religion'],\n",449 " num_rows: 8293\n",450 " })\n",451 " test: Dataset({\n",452 " features: ['ID', 'passage', 'EVENT_Illness', 'EVENT_Accident', 'EVENT_Other', 'CAUSE_Material_Physical', 'CAUSE_Spirits_Gods', 'CAUSE_Witchcraft_Sorcery', 'CAUSE_Rule_Violation_Taboo', 'ACTION_Physical_Material', 'ACTION_Technical_Specialist', 'ACTION_Divination', 'ACTION_Shaman_Medium_Healer', 'ACTION_Priest_High_Religion'],\n",453 " num_rows: 2074\n",454 " })\n",455 "})"456 ]457 },458 "execution_count": 93,459 "metadata": {},460 "output_type": "execute_result"461 }462 ],463 "source": [464 "# subdivide into just passage and outcome\n",465 "df_small = pd.DataFrame()\n",466 "df_small[colNames] = df[col_list]\n",467 "# Flip the lable of \"no_info\"\n",468 "# df_small[[\"EVENT\",\"CAUSE\",\"ACTION\"]] = df_small[[\"EVENT\",\"CAUSE\",\"ACTION\"]].replace({0:1, 1:0})\n",469 "\n",470 "\n",471 "# create train and validation/test sets\n",472 "train_val, test = train_test_split(df_small, test_size=0.2, random_state=10)\n",473 "\n",474 "\n",475 "# Create an NLP friendly dataset\n",476 "Hraf = DatasetDict(\n",477 " {'train':Dataset.from_dict(train_val.to_dict(orient= 'list')),\n",478 " 'test':Dataset.from_dict(test.to_dict(orient= 'list'))})\n",479 "Hraf"480 ]481 },482 {483 "cell_type": "markdown",484 "metadata": {},485 "source": [486 " Show class bias"487 ]488 },489 {490 "cell_type": "code",491 "execution_count": 94,492 "metadata": {},493 "outputs": [494 {495 "name": "stdout",496 "output_type": "stream",497 "text": [498 "BIAS FOR ANSWERING 'PRESENT'\n",499 "Passage Count: 10367\n",500 " Raw Adj.\n",501 "____________________________________________________________\n",502 "\n",503 "EVENT: 0.63\n",504 "\tIllness: 0.41 0.64\n",505 "\tAccident: 0.07 0.1\n",506 "\tOther: 0.26 0.41\n",507 "\n",508 "CAUSE: 0.47\n",509 "\tJust_Happens: 0.02 0.05\n",510 "\tMaterial_Physical: 0.17 0.37\n",511 "\tSpirits_Gods: 0.19 0.4\n",512 "\tWitchcraft_Sorcery: 0.06 0.14\n",513 "\tRule_Violation_Taboo: 0.1 0.22\n",514 "\tOther: 0.05 0.11\n",515 "\n",516 "ACTION: 0.48\n",517 "\tPhysical_Material: 0.32 0.68\n",518 "\tTechnical_Specialist: 0.07 0.14\n",519 "\tDivination: 0.02 0.05\n",520 "\tShaman_Medium_Healer: 0.08 0.16\n",521 "\tPriest_High_Religion: 0.03 0.07\n",522 "\tOther: 0.07 0.15\n"523 ]524 },525 {526 "name": "stderr",527 "output_type": "stream",528 "text": [529 "/var/folders/r8/j6vx966d5rj6srtb4x895kvr0000gp/T/ipykernel_51139/3439163906.py:22: FutureWarning: The behavior of DataFrame concatenation with empty or all-NA entries is deprecated. In a future version, this will no longer exclude empty or all-NA columns when determining the result dtypes. To retain the old behavior, exclude the relevant entries before the concat operation.\n",530 " df_biases = pd.concat([df_biases, pd.DataFrame({\"Class\":[col[0]],\"Raw_Bias\":[proportion]})])\n",531 "/var/folders/r8/j6vx966d5rj6srtb4x895kvr0000gp/T/ipykernel_51139/3439163906.py:28: FutureWarning: The behavior of DataFrame concatenation with empty or all-NA entries is deprecated. In a future version, this will no longer exclude empty or all-NA columns when determining the result dtypes. To retain the old behavior, exclude the relevant entries before the concat operation.\n",532 " df_biases = pd.concat([df_biases, pd.DataFrame({\"Class\":[col[1]],\"Raw_Bias\":[proportion], \"Adj__Bias\":[adj_proportion]})])\n"533 ]534 }535 ],536 "source": [537 "multiCol = list(df.columns)\n",538 "valuecountCol = []\n",539 "for col in multiCol:\n",540 " if col[0] in [\"CULTURE\", \"OTHER\", \"CODER\"] or col[1] in [\"Description\", \"Local_terms\", \"Local_Terms\"]:\n",541 " continue\n",542 " else:\n",543 " valuecountCol.append(col)\n",544 " \n",545 "# set up dataframe for easy saving\n",546 "df_biases = pd.DataFrame(columns=[\"Class\",\"Raw_Bias\",\"Adj__Bias\"])\n",547 "\n",548 "# Get proportions and show table. 'raw' is just number of present divided by total while 'adj' is within main category proportion present divided by total main class present\n",549 "print(\"BIAS FOR ANSWERING \\'PRESENT\\'\")\n",550 "print(\"Passage Count: \", len(df))\n",551 "print(f\"{' '*39}Raw{' '*10}Adj.\") \n",552 "print(f\"{'_'*60}\")\n",553 "for col in valuecountCol:\n",554 " if col[1] == \"No_Info\":\n",555 " proportion = 1-np.mean(df[col])\n",556 " mainCat_proportion = proportion\n",557 " proportion = round(proportion,2)\n",558 " df_biases = pd.concat([df_biases, pd.DataFrame({\"Class\":[col[0]],\"Raw_Bias\":[proportion]})])\n",559 " print(f\"\\n{col[0]}:{(38-len(col[0]))*' '}{proportion}\")\n",560 " else:\n",561 " proportion = np.mean(df[col])\n",562 " adj_proportion = round(proportion / mainCat_proportion,2) # get adjusted proportion within category\n",563 " proportion = round(proportion,2)\n",564 " df_biases = pd.concat([df_biases, pd.DataFrame({\"Class\":[col[1]],\"Raw_Bias\":[proportion], \"Adj__Bias\":[adj_proportion]})])\n",565 " print(f\"\\t{col[1]}:{(30-len(col[1]))*' '}{proportion}{' '*(12-len(str(proportion)))}{adj_proportion}\")\n"566 ]567 },568 {569 "attachments": {},570 "cell_type": "markdown",571 "metadata": {},572 "source": [573 "Make sure the train, validation, and test sets are as biased as our total input data (we want each to match more or less with the total) <br>"574 ]575 },576 {577 "cell_type": "code",578 "execution_count": 95,579 "metadata": {},580 "outputs": [581 {582 "name": "stdout",583 "output_type": "stream",584 "text": [585 " TOTAL train test\n",586 "____________________________________________________________\n",587 "EVENT_Illness: 40.58 40.55 40.69 \n",588 "EVENT_Accident: 6.56 6.48 6.89 \n",589 "EVENT_Other: 26.04 26.19 25.46 \n",590 "CAUSE_Material_Physical: 17.04 17.05 17.02 \n",591 "CAUSE_Spirits_Gods: 18.72 18.45 19.82 \n",592 "CAUSE_Witchcraft_Sorcery: 6.33 6.38 6.12 \n",593 "CAUSE_Rule_Violation_Taboo: 10.19 10.01 10.9 \n",594 "ACTION_Physical_Material: 32.21 32.26 32.02 \n",595 "ACTION_Technical_Specialist: 6.9 7.05 6.27 \n",596 "ACTION_Divination: 2.22 2.18 2.36 \n",597 "ACTION_Shaman_Medium_Healer: 7.72 7.72 7.71 \n",598 "ACTION_Priest_High_Religion: 3.44 3.57 2.94 \n"599 ]600 }601 ],602 "source": [603 "# extract the total proportion\n",604 "def totalProportion(df, col, present=1):\n",605 " value_counts = df[col].value_counts()\n",606 " percentage = round(value_counts[present]/len(df)*100,2)\n",607 " return percentage\n",608 "\n",609 "# extracts percentages per datafaframe\n",610 "def colProportion(Hraf, col):\n",611 " percentage_list = []\n",612 " for dataframe in Hraf.keys():\n",613 " percentage_list += [round(sum(Hraf[dataframe][col]) / (len(Hraf[dataframe]))*100,2)]\n",614 " return percentage_list\n",615 "\n",616 "\n",617 "\n",618 "# print bias per label\n",619 "dataframe_keys= Hraf.keys()\n",620 "labels = [label for label in Hraf['train'].features.keys() if label not in ['ID', 'passage']]\n",621 "header = \" TOTAL\"\n",622 "for key in dataframe_keys:\n",623 " header += f\" {key}\"\n",624 "print(header)\n",625 "print('_'*(len(header)+4))\n",626 "for col in labels:\n",627 " totalPercentage = totalProportion(df_small, col)\n",628 " percentage_list = colProportion(Hraf, col)\n",629 " spacing = 10\n",630 " percentage_str = f\"{totalPercentage}{' '* (spacing-len(str(totalPercentage)))}\"\n",631 " for index, key in enumerate(dataframe_keys):\n",632 " percentage_str += f\"{(len(key)-5)*' '}{percentage_list[index]}{' '* (spacing-len(str(percentage_list[index])))}\"\n",633 " print(f\"{col}:{' ' * (30- len(col))} {percentage_str}\")"634 ]635 },636 {637 "attachments": {},638 "cell_type": "markdown",639 "metadata": {},640 "source": [641 "## Preprocess"642 ]643 },644 {645 "attachments": {},646 "cell_type": "markdown",647 "metadata": {},648 "source": [649 "Create labels for training and preprocessing"650 ]651 },652 {653 "cell_type": "code",654 "execution_count": 96,655 "metadata": {},656 "outputs": [657 {658 "data": {659 "text/plain": [660 "{0: 'EVENT_Illness',\n",661 " 1: 'EVENT_Accident',\n",662 " 2: 'EVENT_Other',\n",663 " 3: 'CAUSE_Material_Physical',\n",664 " 4: 'CAUSE_Spirits_Gods',\n",665 " 5: 'CAUSE_Witchcraft_Sorcery',\n",666 " 6: 'CAUSE_Rule_Violation_Taboo',\n",667 " 7: 'ACTION_Physical_Material',\n",668 " 8: 'ACTION_Technical_Specialist',\n",669 " 9: 'ACTION_Divination',\n",670 " 10: 'ACTION_Shaman_Medium_Healer',\n",671 " 11: 'ACTION_Priest_High_Religion'}"672 ]673 },674 "execution_count": 96,675 "metadata": {},676 "output_type": "execute_result"677 }678 ],679 "source": [680 "\n",681 "labels = [label for label in Hraf['train'].features.keys() if label not in ['ID', 'passage']]\n",682 "id2label = {idx:label for idx, label in enumerate(labels)}\n",683 "label2id = {label:idx for idx, label in enumerate(labels)}\n",684 "id2label"685 ]686 },687 {688 "attachments": {},689 "cell_type": "markdown",690 "metadata": {},691 "source": [692 "load a DistilBERT tokenizer to preprocess the text field: <br>"693 ]694 },695 {696 "attachments": {},697 "cell_type": "markdown",698 "metadata": {},699 "source": [700 "Create a preprocessing function to tokenize text and truncate sequences to be no longer than DistilBERT’s maximum input length:<br>\n",701 "Guidelines were followed from NielsRogge found <a href= \"https://github.com/NielsRogge/Transformers-Tutorials/blob/master/BERT/Fine_tuning_BERT_(and_friends)_for_multi_label_text_classification.ipynb\"> here </a>"702 ]703 },704 {705 "cell_type": "code",706 "execution_count": 97,707 "metadata": {},708 "outputs": [],709 "source": [710 "from transformers import AutoTokenizer\n",711 "import numpy as np\n",712 "\n",713 "\n",714 "tokenizer = AutoTokenizer.from_pretrained('distilbert-base-uncased')\n",715 "\n",716 "def preprocess_data(examples):\n",717 " # take a batch of texts\n",718 " text = examples[\"passage\"]\n",719 " # encode them\n",720 " encoding = tokenizer(text, max_length=512, truncation=True) #max length for BERT is 512\n",721 " # add labels\n",722 " labels_batch = {k: examples[k] for k in examples.keys() if k in labels}\n",723 " # create numpy array of shape (batch_size, num_labels)\n",724 " labels_matrix = np.zeros((len(text), len(labels)))\n",725 " # fill numpy array\n",726 " for idx, label in enumerate(labels):\n",727 " labels_matrix[:, idx] = labels_batch[label]\n",728 "\n",729 " encoding[\"labels\"] = labels_matrix.tolist()\n",730 "\n",731 " return encoding"732 ]733 },734 {735 "attachments": {},736 "cell_type": "markdown",737 "metadata": {},738 "source": [739 "To apply the preprocessing function over the entire dataset, use 🤗 Datasets map function. You can speed up map by setting batched=True to process multiple elements of the dataset at once:"740 ]741 },742 {743 "cell_type": "code",744 "execution_count": 98,745 "metadata": {},746 "outputs": [747 {748 "data": {749 "application/vnd.jupyter.widget-view+json": {750 "model_id": "7a8c7ac9531b4321aa70c9a45eb53ac0",751 "version_major": 2,752 "version_minor": 0753 },754 "text/plain": [755 "Map: 0%| | 0/8293 [00:00<?, ? examples/s]"756 ]757 },758 "metadata": {},759 "output_type": "display_data"760 },761 {762 "data": {763 "application/vnd.jupyter.widget-view+json": {764 "model_id": "c5e5f44f009a436d9aa07fcb66d1bc0b",765 "version_major": 2,766 "version_minor": 0767 },768 "text/plain": [769 "Map: 0%| | 0/2074 [00:00<?, ? examples/s]"770 ]771 },772 "metadata": {},773 "output_type": "display_data"774 }775 ],776 "source": [777 "# Tokenize data, remove all columns and give new ones\n",778 "tokenized_Hraf = Hraf.map(preprocess_data, batched=True, remove_columns=Hraf['train'].column_names)"779 ]780 },781 {782 "cell_type": "code",783 "execution_count": 99,784 "metadata": {},785 "outputs": [786 {787 "data": {788 "text/plain": [789 "DatasetDict({\n",790 " train: Dataset({\n",791 " features: ['input_ids', 'attention_mask', 'labels'],\n",792 " num_rows: 8293\n",793 " })\n",794 " test: Dataset({\n",795 " features: ['input_ids', 'attention_mask', 'labels'],\n",796 " num_rows: 2074\n",797 " })\n",798 "})"799 ]800 },801 "execution_count": 99,802 "metadata": {},803 "output_type": "execute_result"804 }805 ],806 "source": [807 "# Set tokenized passages to PyTorch Tensor\n",808 "tokenized_Hraf.set_format(\"torch\")\n",809 "tokenized_Hraf"810 ]811 },812 {813 "cell_type": "code",814 "execution_count": 100,815 "metadata": {},816 "outputs": [817 {818 "name": "stdout",819 "output_type": "stream",820 "text": [821 "dict_keys(['input_ids', 'attention_mask', 'labels'])\n",822 "[CLS] among the ornaments there are also little copper bells / dutch klokjes and belletjes, both words meaning “ little bells ” /, which are cast by moriers and imported ( dio - dio, bangkoela ). in former times they could be worn only by men who had already slain several enemies and by priestesses ; before they were put on, they were counted off on the wearer from 1 to 7. nowadays young people often go about with them in order to attract attention to themselves through the tinkling. men hang the little bell / klokje / on the band on which they wear their sword, so that while walking it continuously swings against their legs and tinkles. the priestess has it hanging from her belt. she makes use of it on various occasions ( see index under “ klokje ” ). in connection with agriculture a bell is tinkled now and then when people start performing the first work of clearing ( mombakati ) ; this is said to serve to render harmless sounds prophesying evil ; at the time the rice is supposed to sprout, there is tinkling so that the ears will come out at the same time ; and at the beginning of the harvest a bell is sounded next to the bound - together stools ( pesoea ), in order to summon the field spirits ( lamoa nawoe ) ( onda ’ e ). [SEP]\n",823 "tensor([0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0.])\n"824 ]825 },826 {827 "data": {828 "text/plain": [829 "[]"830 ]831 },832 "execution_count": 100,833 "metadata": {},834 "output_type": "execute_result"835 }836 ],837 "source": [838 "example = tokenized_Hraf['train'][1]\n",839 "print(example.keys())\n",840 "print(tokenizer.decode(example['input_ids']))\n",841 "print(example['labels'])\n",842 "[id2label[idx] for idx, label in enumerate(example['labels']) if label == 1.0]\n"843 ]844 },845 {846 "cell_type": "code",847 "execution_count": 101,848 "metadata": {},849 "outputs": [850 {851 "name": "stdout",852 "output_type": "stream",853 "text": [854 "Number Truncated: 160\n",855 "Percentage Truncated: 1.9%\n",856 "[207, 322, 324, 446, 458, 484, 702, 721, 745, 801, 805, 962, 1011, 1191, 1194, 1219, 1231, 1295, 1314, 1370, 1459, 1470, 1574, 1614, 1680, 1686, 1861, 1901, 1935, 1985, 2005, 2019, 2026, 2031, 2034, 2051, 2130, 2145, 2160, 2225, 2279, 2305, 2333, 2393, 2398, 2445, 2459, 2514, 2546, 2557, 2562, 2598, 2601, 2655, 2656, 2661, 2663, 2725, 2788, 2851, 3067, 3101, 3133, 3274, 3313, 3315, 3537, 3557, 3670, 3676, 3800, 3818, 3843, 3929, 3988, 4066, 4097, 4173, 4175, 4213, 4264, 4280, 4288, 4363, 4429, 4430, 4514, 4540, 4590, 4624, 4801, 4816, 4876, 4889, 4895, 4912, 5097, 5149, 5241, 5349, 5368, 5474, 5509, 5541, 5632, 5767, 5786, 5792, 5800, 5855, 5863, 5867, 5970, 5972, 5977, 6053, 6060, 6069, 6108, 6139, 6172, 6190, 6366, 6410, 6456, 6506, 6514, 6634, 6641, 6656, 6666, 6708, 6767, 6790, 6820, 6859, 6877, 6951, 6970, 6990, 7062, 7080, 7097, 7111, 7236, 7368, 7630, 7634, 7654, 7704, 7795, 7829, 7896, 7904, 7907, 8060, 8154, 8157, 8179, 8251]\n"857 ]858 }859 ],860 "source": [861 "# Show number of passages longer than 512 tokens (and therefore truncated)\n",862 "sequence_i = []\n",863 "for i, tx in enumerate(tokenized_Hraf['train']):\n",864 " if len(tx['input_ids']) == 512:\n",865 " sequence_i.append(i)\n",866 "print('Number Truncated: ', len(sequence_i))\n",867 "print(f'Percentage Truncated: {round(len(sequence_i)/len(tokenized_Hraf[\"train\"])*100,1)}%')\n",868 "print(sequence_i)"869 ]870 },871 {872 "attachments": {},873 "cell_type": "markdown",874 "metadata": {},875 "source": [876 "Now create a batch of examples using <a href=\"https://huggingface.co/docs/transformers/v4.29.0/en/main_classes/data_collator#transformers.DataCollatorWithPadding\"> DataCollatorWithPadding</a>. It’s more efficient to dynamically pad the sentences to the longest length in a batch during collation, instead of padding the whole dataset to the maximum length."877 ]878 },879 {880 "cell_type": "markdown",881 "metadata": {},882 "source": [883 "### Create Splits"884 ]885 },886 {887 "cell_type": "markdown",888 "metadata": {},889 "source": [890 " Stratification using multilabels is a difficult process as the number of unique bins of stratification increases exponentially by the number of labels (see more info and potential ways to conduct multilabel sttratification sampling <a href=\"https://dl.acm.org/doi/10.5555/2034161.2034172\"> HERE </a>). We will currently disregard focusing on stratification of all the labels/classifications and just use a single label for stratification. Currently, this is still giving decent splits that do not deviate far from the true proportion or between n_splits. Still, one should check the proportional deviation of each label to make sure"891 ]892 },893 {894 "cell_type": "code",895 "execution_count": 102,896 "metadata": {},897 "outputs": [898 {899 "data": {900 "text/html": [901 "<div>\n",902 "<style scoped>\n",903 " .dataframe tbody tr th:only-of-type {\n",904 " vertical-align: middle;\n",905 " }\n",906 "\n",907 " .dataframe tbody tr th {\n",908 " vertical-align: top;\n",909 " }\n",910 "\n",911 " .dataframe thead th {\n",912 " text-align: right;\n",913 " }\n",914 "</style>\n",915 "<table border=\"1\" class=\"dataframe\">\n",916 " <thead>\n",917 " <tr style=\"text-align: right;\">\n",918 " <th></th>\n",919 " <th>EVENT_Illness</th>\n",920 " <th>EVENT_Accident</th>\n",921 " <th>EVENT_Other</th>\n",922 " <th>CAUSE_Material_Physical</th>\n",923 " <th>CAUSE_Spirits_Gods</th>\n",924 " <th>CAUSE_Witchcraft_Sorcery</th>\n",925 " <th>CAUSE_Rule_Violation_Taboo</th>\n",926 " <th>ACTION_Physical_Material</th>\n",927 " <th>ACTION_Technical_Specialist</th>\n",928 " <th>ACTION_Divination</th>\n",929 " <th>ACTION_Shaman_Medium_Healer</th>\n",930 " <th>ACTION_Priest_High_Religion</th>\n",931 " </tr>\n",932 " </thead>\n",933 " <tbody>\n",934 " <tr>\n",935 " <th>Fold 1</th>\n",936 " <td>0.41</td>\n",937 " <td>0.07</td>\n",938 " <td>0.26</td>\n",939 " <td>0.17</td>\n",940 " <td>0.19</td>\n",941 " <td>0.07</td>\n",942 " <td>0.1</td>\n",943 " <td>0.33</td>\n",944 " <td>0.07</td>\n",945 " <td>0.02</td>\n",946 " <td>0.08</td>\n",947 " <td>0.04</td>\n",948 " </tr>\n",949 " <tr>\n",950 " <th>Fold 2</th>\n",951 " <td>0.41</td>\n",952 " <td>0.06</td>\n",953 " <td>0.26</td>\n",954 " <td>0.17</td>\n",955 " <td>0.18</td>\n",956 " <td>0.06</td>\n",957 " <td>0.1</td>\n",958 " <td>0.32</td>\n",959 " <td>0.07</td>\n",960 " <td>0.02</td>\n",961 " <td>0.08</td>\n",962 " <td>0.04</td>\n",963 " </tr>\n",964 " <tr>\n",965 " <th>Fold 3</th>\n",966 " <td>0.40</td>\n",967 " <td>0.07</td>\n",968 " <td>0.26</td>\n",969 " <td>0.17</td>\n",970 " <td>0.18</td>\n",971 " <td>0.06</td>\n",972 " <td>0.1</td>\n",973 " <td>0.32</td>\n",974 " <td>0.07</td>\n",975 " <td>0.02</td>\n",976 " <td>0.08</td>\n",977 " <td>0.04</td>\n",978 " </tr>\n",979 " <tr>\n",980 " <th>Fold 4</th>\n",981 " <td>0.40</td>\n",982 " <td>0.06</td>\n",983 " <td>0.26</td>\n",984 " <td>0.17</td>\n",985 " <td>0.19</td>\n",986 " <td>0.07</td>\n",987 " <td>0.1</td>\n",988 " <td>0.32</td>\n",989 " <td>0.07</td>\n",990 " <td>0.02</td>\n",991 " <td>0.08</td>\n",992 " <td>0.04</td>\n",993 " </tr>\n",994 " <tr>\n",995 " <th>Fold 5</th>\n",996 " <td>0.41</td>\n",997 " <td>0.06</td>\n",998 " <td>0.26</td>\n",999 " <td>0.17</td>\n",1000 " <td>0.18</td>\n",1001 " <td>0.06</td>\n",1002 " <td>0.1</td>\n",1003 " <td>0.33</td>\n",1004 " <td>0.07</td>\n",1005 " <td>0.02</td>\n",1006 " <td>0.08</td>\n",1007 " <td>0.04</td>\n",1008 " </tr>\n",1009 " </tbody>\n",1010 "</table>\n",1011 "</div>"1012 ],1013 "text/plain": [1014 " EVENT_Illness EVENT_Accident EVENT_Other CAUSE_Material_Physical \\\n",1015 "Fold 1 0.41 0.07 0.26 0.17 \n",1016 "Fold 2 0.41 0.06 0.26 0.17 \n",1017 "Fold 3 0.40 0.07 0.26 0.17 \n",1018 "Fold 4 0.40 0.06 0.26 0.17 \n",1019 "Fold 5 0.41 0.06 0.26 0.17 \n",1020 "\n",1021 " CAUSE_Spirits_Gods CAUSE_Witchcraft_Sorcery \\\n",1022 "Fold 1 0.19 0.07 \n",1023 "Fold 2 0.18 0.06 \n",1024 "Fold 3 0.18 0.06 \n",1025 "Fold 4 0.19 0.07 \n",1026 "Fold 5 0.18 0.06 \n",1027 "\n",1028 " CAUSE_Rule_Violation_Taboo ACTION_Physical_Material \\\n",1029 "Fold 1 0.1 0.33 \n",1030 "Fold 2 0.1 0.32 \n",1031 "Fold 3 0.1 0.32 \n",1032 "Fold 4 0.1 0.32 \n",1033 "Fold 5 0.1 0.33 \n",1034 "\n",1035 " ACTION_Technical_Specialist ACTION_Divination \\\n",1036 "Fold 1 0.07 0.02 \n",1037 "Fold 2 0.07 0.02 \n",1038 "Fold 3 0.07 0.02 \n",1039 "Fold 4 0.07 0.02 \n",1040 "Fold 5 0.07 0.02 \n",1041 "\n",1042 " ACTION_Shaman_Medium_Healer ACTION_Priest_High_Religion \n",1043 "Fold 1 0.08 0.04 \n",1044 "Fold 2 0.08 0.04 \n",1045 "Fold 3 0.08 0.04 \n",1046 "Fold 4 0.08 0.04 \n",1047 "Fold 5 0.08 0.04 "1048 ]1049 },1050 "execution_count": 102,1051 "metadata": {},1052 "output_type": "execute_result"1053 }1054 ],1055 "source": [1056 "# Splitting\n",1057 "from sklearn.model_selection import StratifiedKFold\n",1058 "fold_n =5\n",1059 "\n",1060 "# folds = StratifiedKFold(n_splits=5)\n",1061 "folds = StratifiedKFold(n_splits=fold_n, shuffle= True, random_state=10)\n",1062 "cols = Hraf['train'].column_names\n",1063 "splits = folds.split(np.zeros(Hraf['train'].num_rows), Hraf['train'][cols[-1]])\n",1064 "# preconstruct dataframe to show\n",1065 "fold_str = [\"Fold \"+str(x) for x in range(1,fold_n+1)]\n",1066 "df_foldPerc = pd.DataFrame(data=np.zeros((fold_n,len(labels))),columns=labels, index=fold_str)\n",1067 "\n",1068 "train_list = []\n",1069 "val_list = []\n",1070 "\n",1071 "for fold, (train_idxs, val_idxs) in enumerate(splits, start=1):\n",1072 " train_list += [train_idxs]\n",1073 " val_list += [val_idxs]\n",1074 " train_hub = Hraf['train'][train_idxs]\n",1075 "\n",1076 " df_foldPerc.iloc[fold-1] = [np.round(np.mean(train_hub[col]),2) for col in cols[2:]]\n",1077 " \n",1078 "df_foldPerc"1079 ]1080 },1081 {1082 "cell_type": "markdown",1083 "metadata": {},1084 "source": [1085 "### Save Paritioned Datasets"1086 ]1087 },1088 {1089 "cell_type": "code",1090 "execution_count": 17,1091 "metadata": {},1092 "outputs": [1093 {1094 "name": "stdout",1095 "output_type": "stream",1096 "text": [1097 "8293 Rows for 'train' succesfully saved to /Users/ericchantland/Library/CloudStorage/Dropbox/MEM-DEV-LAB-Current/2023-eHRAF-Misf/HRAF-Misf-NaturalLanguageProcessing/HRAF_NLP/HRAF_MultiLabel_SubClasses_Kfolds/Datasets/train_dataset.json\n",1098 "2074 Rows for 'test' succesfully saved to /Users/ericchantland/Library/CloudStorage/Dropbox/MEM-DEV-LAB-Current/2023-eHRAF-Misf/HRAF-Misf-NaturalLanguageProcessing/HRAF_NLP/HRAF_MultiLabel_SubClasses_Kfolds/Datasets/test_dataset.json\n"1099 ]1100 }1101 ],1102 "source": [1103 "# # Save datasets for later inference (SKIP IF YOU DO NOT WANT TO OVERWRITE DATASET FILES)\n",1104 "\n",1105 "# def make_dir(path):\n",1106 "# import os\n",1107 "# # Check whether the specified path exists or not\n",1108 "# isExist = os.path.exists(path)\n",1109 "# if not isExist:\n",1110 "# # Create a new directory because it does not exist\n",1111 "# os.makedirs(path)\n",1112 "\n",1113 "# # make folder if it does not exist yet\n",1114 "# path_datasets = os.getcwd() + '/Datasets'\n",1115 "# make_dir(path_datasets)\n",1116 "# # save to Json\n",1117 "# for key in Hraf.keys():\n",1118 "# Hraf_dict = Hraf[key].to_dict()\n",1119 "# file_path = f\"{path_datasets}/{key}_dataset.json\"\n",1120 "# with open(file_path, \"w\") as outfile:\n",1121 "# json.dump(Hraf_dict, outfile)\n",1122 "# print(len(Hraf_dict['ID']), f\"Rows for \\'{key}\\' succesfully saved to {file_path}\")"1123 ]1124 },1125 {1126 "attachments": {},1127 "cell_type": "markdown",1128 "metadata": {},1129 "source": [1130 "## Evaluate"1131 ]1132 },1133 {1134 "attachments": {},1135 "cell_type": "markdown",1136 "metadata": {},1137 "source": [1138 "Obtain F1 score for evaluation"1139 ]1140 },1141 {1142 "cell_type": "code",1143 "execution_count": 103,1144 "metadata": {},1145 "outputs": [],1146 "source": [1147 "from sklearn.metrics import f1_score, roc_auc_score, accuracy_score\n",1148 "from transformers import EvalPrediction, TrainerCallback\n",1149 "import torch\n",1150 "\n",1151 "# Get Metric performance\n",1152 "# source: https://jesusleal.io/2021/04/21/Longformer-multilabel-classification/\n",1153 "def multi_label_metrics(predictions, labels, threshold=0.5):\n",1154 " # first, apply sigmoid on predictions which are of shape (batch_size, num_labels)\n",1155 " sigmoid = torch.nn.Sigmoid()\n",1156 " probs = sigmoid(torch.Tensor(predictions))\n",1157 " # next, use threshold to turn them into integer predictions\n",1158 " y_pred = np.zeros(probs.shape)\n",1159 " y_pred[np.where(probs >= threshold)] = 1\n",1160 " # finally, compute metrics\n",1161 " y_true = labels\n",1162 " f1_micro_average = f1_score(y_true=y_true, y_pred=y_pred, average='micro')\n",1163 " roc_auc = roc_auc_score(y_true, y_pred, average = 'micro')\n",1164 " accuracy = accuracy_score(y_true, y_pred)\n",1165 " # return as dictionary\n",1166 " metrics = {'f1': f1_micro_average,\n",1167 " 'roc_auc': roc_auc,\n",1168 " 'accuracy': accuracy}\n",1169 " return metrics\n",1170 "\n",1171 "# Compute evaluation\n",1172 "def compute_metrics(p: EvalPrediction):\n",1173 " preds = p.predictions[0] if isinstance(p.predictions, \n",1174 " tuple) else p.predictions\n",1175 " result = multi_label_metrics(\n",1176 " predictions=preds, \n",1177 " labels=p.label_ids)\n",1178 " return result\n",1179 "\n"1180 ]1181 },1182 {1183 "attachments": {},1184 "cell_type": "markdown",1185 "metadata": {},1186 "source": [1187 "\n",1188 "## Train\n",1189 "Before you start training your model, create a map of the expected ids to their labels with id2label and label2id:"1190 ]1191 },1192 {1193 "cell_type": "code",1194 "execution_count": 104,1195 "metadata": {},1196 "outputs": [1197 {1198 "name": "stderr",1199 "output_type": "stream",1200 "text": [