{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Introduction\n**This notebook introduced how to solve an imbalanced text classification problem with LSTM networks and word embedding.**","metadata":{}},{"cell_type":"markdown","source":"Import some required libraries.","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport tensorflow as tf\nimport matplotlib.pyplot as plt\n%matplotlib inline\n\nimport gc\nimport sys\n\nfrom tqdm.notebook import tqdm\ntqdm().pandas()\npd.set_option('display.max_colwidth', None)\n\n# Set seed for experiment reproducibility\nseed = 1024\ntf.random.set_seed(seed)\nnp.random.seed(seed)\n\ndef print_size(var):  \n    print('%.2fMB' % (sys.getsizeof(var)/1024/1024))","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Load the train and test dataset.","metadata":{}},{"cell_type":"code","source":"train_data = pd.read_csv('/kaggle/input/quora-insincere-questions-classification/train.csv')\ntest_data = pd.read_csv('/kaggle/input/quora-insincere-questions-classification/test.csv')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"It's not necessary using entire dataset to train if you just run it quickly.","metadata":{}},{"cell_type":"code","source":"# train_data = train_data[0:100000]\n# test_data = test_data[0:10000]","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Let's see what's in the dataset and print the first 5 rows in train data.","metadata":{}},{"cell_type":"code","source":"train_data.head()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Let's see how imbalanced the dataset is.","metadata":{}},{"cell_type":"code","source":"negative, positive = np.bincount(train_data['target'])\ntotal = negative + positive\nprint('total: {}    positive: {} ({:.2f}% of total)'.format(total, positive, 100 * positive / total))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Word vectorizing. converting words into numbers so that can be fed into neural network.","metadata":{}},{"cell_type":"code","source":"import re\n\ndef clean_tag(text):\n    if '[math]' in text:\n        text = re.sub('\\[math\\].*?math\\]', '[formula]', text) #replacing with [formuala]\n\n    if 'http' in text or 'www' in text:\n        text = re.sub('(?:(?:https?|ftp):\\/\\/)?[\\w/\\-?=%.]+\\.[\\w/\\-?=%.]+', '[url]', text) #replacing with [url]\n    return text\n\ncontraction_mapping = {\"We'd\": \"We had\", \"That'd\": \"That had\", \"AREN'T\": \"Are not\", \"HADN'T\": \"Had not\", \"Could've\": \"Could have\", \"LeT's\": \"Let us\", \"How'll\": \"How will\", \"They'll\": \"They will\", \"DOESN'T\": \"Does not\", \"HE'S\": \"He has\", \"O'Clock\": \"Of the clock\", \"Who'll\": \"Who will\", \"What'S\": \"What is\", \"Ain't\": \"Am not\", \"WEREN'T\": \"Were not\", \"Y'all\": \"You all\", \"Y'ALL\": \"You all\", \"Here's\": \"Here is\", \"It'd\": \"It had\", \"Should've\": \"Should have\", \"I'M\": \"I am\", \"ISN'T\": \"Is not\", \"Would've\": \"Would have\", \"He'll\": \"He will\", \"DON'T\": \"Do not\", \"She'd\": \"She had\", \"WOULDN'T\": \"Would not\", \"She'll\": \"She will\", \"IT's\": \"It is\", \"There'd\": \"There had\", \"It'll\": \"It will\", \"You'll\": \"You will\", \"He'd\": \"He had\", \"What'll\": \"What will\", \"Ma'am\": \"Madam\", \"CAN'T\": \"Can not\", \"THAT'S\": \"That is\", \"You've\": \"You have\", \"She's\": \"She is\", \"Weren't\": \"Were not\", \"They've\": \"They have\", \"Couldn't\": \"Could not\", \"When's\": \"When is\", \"Haven't\": \"Have not\", \"We'll\": \"We will\", \"That's\": \"That is\", \"We're\": \"We are\", \"They're\": \"They' are\", \"You'd\": \"You would\", \"How'd\": \"How did\", \"What're\": \"What are\", \"Hasn't\": \"Has not\", \"Wasn't\": \"Was not\", \"Won't\": \"Will not\", \"There's\": \"There is\", \"Didn't\": \"Did not\", \"Doesn't\": \"Does not\", \"You're\": \"You are\", \"He's\": \"He is\", \"SO's\": \"So is\", \"We've\": \"We have\", \"Who's\": \"Who is\", \"Wouldn't\": \"Would not\", \"Why's\": \"Why is\", \"WHO's\": \"Who is\", \"Let's\": \"Let us\", \"How's\": \"How is\", \"Can't\": \"Can not\", \"Where's\": \"Where is\", \"They'd\": \"They had\", \"Don't\": \"Do not\", \"Shouldn't\":\"Should not\", \"Aren't\":\"Are not\", \"ain't\": \"is not\", \"What's\": \"What is\", \"It's\": \"It is\", \"Isn't\":\"Is not\", \"aren't\": \"are not\",\"can't\": \"cannot\", \"'cause\": \"because\", \"could've\": \"could have\", \"couldn't\": \"could not\", \"didn't\": \"did not\",  \"doesn't\": \"does not\", \"don't\": \"do not\", \"hadn't\": \"had not\", \"hasn't\": \"has not\", \"haven't\": \"have not\", \"he'd\": \"he would\",\"he'll\": \"he will\", \"he's\": \"he is\", \"how'd\": \"how did\", \"how'd'y\": \"how do you\", \"how'll\": \"how will\", \"how's\": \"how is\",  \"I'd\": \"I would\", \"I'd've\": \"I would have\", \"I'll\": \"I will\", \"I'll've\": \"I will have\",\"I'm\": \"I am\", \"I've\": \"I have\", \"i'd\": \"i would\", \"i'd've\": \"i would have\", \"i'll\": \"i will\",  \"i'll've\": \"i will have\",\"i'm\": \"i am\", \"i've\": \"i have\", \"isn't\": \"is not\", \"it'd\": \"it would\", \"it'd've\": \"it would have\", \"it'll\": \"it will\", \"it'll've\": \"it will have\",\"it's\": \"it is\", \"let's\": \"let us\", \"ma'am\": \"madam\", \"mayn't\": \"may not\", \"might've\": \"might have\",\"mightn't\": \"might not\",\"mightn't've\": \"might not have\", \"must've\": \"must have\", \"mustn't\": \"must not\", \"mustn't've\": \"must not have\", \"needn't\": \"need not\", \"needn't've\": \"need not have\",\"o'clock\": \"of the clock\", \"oughtn't\": \"ought not\", \"oughtn't've\": \"ought not have\", \"shan't\": \"shall not\", \"sha'n't\": \"shall not\", \"shan't've\": \"shall not have\", \"she'd\": \"she would\", \"she'd've\": \"she would have\", \"she'll\": \"she will\", \"she'll've\": \"she will have\", \"she's\": \"she is\", \"should've\": \"should have\", \"shouldn't\": \"should not\", \"shouldn't've\": \"should not have\", \"so've\": \"so have\",\"so's\": \"so as\", \"this's\": \"this is\",\"that'd\": \"that would\", \"that'd've\": \"that would have\", \"that's\": \"that is\", \"there'd\": \"there would\", \"there'd've\": \"there would have\", \"there's\": \"there is\", \"here's\": \"here is\",\"they'd\": \"they would\", \"they'd've\": \"they would have\", \"they'll\": \"they will\", \"they'll've\": \"they will have\", \"they're\": \"they are\", \"they've\": \"they have\", \"to've\": \"to have\", \"wasn't\": \"was not\", \"we'd\": \"we would\", \"we'd've\": \"we would have\", \"we'll\": \"we will\", \"we'll've\": \"we will have\", \"we're\": \"we are\", \"we've\": \"we have\", \"weren't\": \"were not\", \"what'll\": \"what will\", \"what'll've\": \"what will have\", \"what're\": \"what are\",  \"what's\": \"what is\", \"what've\": \"what have\", \"when's\": \"when is\", \"when've\": \"when have\", \"where'd\": \"where did\", \"where's\": \"where is\", \"where've\": \"where have\", \"who'll\": \"who will\", \"who'll've\": \"who will have\", \"who's\": \"who is\", \"who've\": \"who have\", \"why's\": \"why is\", \"why've\": \"why have\", \"will've\": \"will have\", \"won't\": \"will not\", \"won't've\": \"will not have\", \"would've\": \"would have\", \"wouldn't\": \"would not\", \"wouldn't've\": \"would not have\", \"y'all\": \"you all\", \"y'all'd\": \"you all would\",\"y'all'd've\": \"you all would have\",\"y'all're\": \"you all are\",\"y'all've\": \"you all have\",\"you'd\": \"you would\", \"you'd've\": \"you would have\", \"you'll\": \"you will\", \"you'll've\": \"you will have\", \"you're\": \"you are\", \"you've\": \"you have\" }\n\ndef clean_contractions(text):\n    specials = [\"’\", \"‘\", \"´\", \"`\"]\n    for s in specials:\n        text = text.replace(s, \"'\")\n    \n    text = ' '.join([contraction_mapping[t] if t in contraction_mapping else t for t in text.split(\" \")])\n    return text\n\npuncts = [\",\",\".\",'\"',\":\",\")\",\"(\",\"-\",\"!\",\"?\",\"|\",\";\",\"'\",\"$\",\"&\",\"/\",\"[\",\"]\",\">\",\"%\",\"=\",\"#\",\"*\",\"+\",\"\\\\\",\"•\",\"~\",\"@\",\"£\",\"·\",\"_\",\"{\",\"}\",\"©\",\"^\",\"®\",\"`\",\"<\",\"→\",\"°\",\"€\",\"™\",\"›\",\"♥\",\"←\",\"×\",\"§\",\"″\",\"′\",\"█\",\"…\",\"“\",\"★\",\"”\",\"–\",\"●\",\"►\",\"−\",\"¢\",\"¬\",\"░\",\"¡\",\"¶\",\"↑\",\"±\",\"¿\",\"▾\",\"═\",\"¦\",\"║\",\"―\",\"¥\",\"▓\",\"—\",\"‹\",\"─\",\"▒\",\"：\",\"⊕\",\"▼\",\"▪\",\"†\",\"■\",\"’\",\"▀\",\"¨\",\"▄\",\"♫\",\"☆\",\"¯\",\"♦\",\"¤\",\"▲\",\"¸\",\"⋅\",\"‘\",\"∞\",\"∙\",\"）\",\"↓\",\"、\",\"│\",\"（\",\"»\",\"，\",\"♪\",\"╩\",\"╚\",\"・\",\"╦\",\"╣\",\"╔\",\"╗\",\"▬\",\"❤\",\"≤\",\"‡\",\"√\",\"◄\",\"━\",\"⇒\",\"▶\",\"≥\",\"╝\",\"♡\",\"◊\",\"。\",\"✈\",\"≡\",\"☺\",\"✔\",\"↵\",\"≈\",\"✓\",\"♣\",\"☎\",\"℃\",\"◦\",\"└\",\"‟\",\"～\",\"！\",\"○\",\"◆\",\"№\",\"♠\",\"▌\",\"✿\",\"▸\",\"⁄\",\"□\",\"❖\",\"✦\",\"．\",\"÷\",\"｜\",\"┃\",\"／\",\"￥\",\"╠\",\"↩\",\"✭\",\"▐\",\"☼\",\"☻\",\"┐\",\"├\",\"«\",\"∼\",\"┌\",\"℉\",\"☮\",\"฿\",\"≦\",\"♬\",\"✧\",\"〉\",\"－\",\"⌂\",\"✖\",\"･\",\"◕\",\"※\",\"‖\",\"◀\",\"‰\",\"\\x97\",\"↺\",\"∆\",\"┘\",\"┬\",\"╬\",\"،\",\"⌘\",\"⊂\",\"＞\",\"〈\",\"⎙\",\"？\",\"☠\",\"⇐\",\"▫\",\"∗\",\"∈\",\"≠\",\"♀\",\"♔\",\"˚\",\"℗\",\"┗\",\"＊\",\"┼\",\"❀\",\"＆\",\"∩\",\"♂\",\"‿\",\"∑\",\"‣\",\"➜\",\"┛\",\"⇓\",\"☯\",\"⊖\",\"☀\",\"┳\",\"；\",\"∇\",\"⇑\",\"✰\",\"◇\",\"♯\",\"☞\",\"´\",\"↔\",\"┏\",\"｡\",\"◘\",\"∂\",\"✌\",\"♭\",\"┣\",\"┴\",\"┓\",\"✨\",\"\\xa0\",\"˜\",\"❥\",\"┫\",\"℠\",\"✒\",\"［\",\"∫\",\"\\x93\",\"≧\",\"］\",\"\\x94\",\"∀\",\"♛\",\"\\x96\",\"∨\",\"◎\",\"↻\",\"⇩\",\"＜\",\"≫\",\"✩\",\"✪\",\"♕\",\"؟\",\"₤\",\"☛\",\"╮\",\"␊\",\"＋\",\"┈\",\"％\",\"╋\",\"▽\",\"⇨\",\"┻\",\"⊗\",\"￡\",\"।\",\"▂\",\"✯\",\"▇\",\"＿\",\"➤\",\"✞\",\"＝\",\"▷\",\"△\",\"◙\",\"▅\",\"✝\",\"∧\",\"␉\",\"☭\",\"┊\",\"╯\",\"☾\",\"➔\",\"∴\",\"\\x92\",\"▃\",\"↳\",\"＾\",\"׳\",\"➢\",\"╭\",\"➡\",\"＠\",\"⊙\",\"☢\",\"˝\",\"∏\",\"„\",\"∥\",\"❝\",\"☐\",\"▆\",\"╱\",\"⋙\",\"๏\",\"☁\",\"⇔\",\"▔\",\"\\x91\",\"➚\",\"◡\",\"╰\",\"\\x85\",\"♢\",\"˙\",\"۞\",\"✘\",\"✮\",\"☑\",\"⋆\",\"ⓘ\",\"❒\",\"☣\",\"✉\",\"⌊\",\"➠\",\"∣\",\"❑\",\"◢\",\"ⓒ\",\"\\x80\",\"〒\",\"∕\",\"▮\",\"⦿\",\"✫\",\"✚\",\"⋯\",\"♩\",\"☂\",\"❞\",\"‗\",\"܂\",\"☜\",\"‾\",\"✜\",\"╲\",\"∘\",\"⟩\",\"＼\",\"⟨\",\"·\",\"✗\",\"♚\",\"∅\",\"ⓔ\",\"◣\",\"͡\",\"‛\",\"❦\",\"◠\",\"✄\",\"❄\",\"∃\",\"␣\",\"≪\",\"｢\",\"≅\",\"◯\",\"☽\",\"∎\",\"｣\",\"❧\",\"̅\",\"ⓐ\",\"↘\",\"⚓\",\"▣\",\"˘\",\"∪\",\"⇢\",\"✍\",\"⊥\",\"＃\",\"⎯\",\"↠\",\"۩\",\"☰\",\"◥\",\"⊆\",\"✽\",\"⚡\",\"↪\",\"❁\",\"☹\",\"◼\",\"☃\",\"◤\",\"❏\",\"ⓢ\",\"⊱\",\"➝\",\"̣\",\"✡\",\"∠\",\"｀\",\"▴\",\"┤\",\"∝\",\"♏\",\"ⓐ\",\"✎\",\";\",\"␤\",\"＇\",\"❣\",\"✂\",\"✤\",\"ⓞ\",\"☪\",\"✴\",\"⌒\",\"˛\",\"♒\",\"＄\",\"✶\",\"▻\",\"ⓔ\",\"◌\",\"◈\",\"❚\",\"❂\",\"￦\",\"◉\",\"╜\",\"̃\",\"✱\",\"╖\",\"❉\",\"ⓡ\",\"↗\",\"ⓣ\",\"♻\",\"➽\",\"׀\",\"✲\",\"✬\",\"☉\",\"▉\",\"≒\",\"☥\",\"⌐\",\"♨\",\"✕\",\"ⓝ\",\"⊰\",\"❘\",\"＂\",\"⇧\",\"̵\",\"➪\",\"▁\",\"▏\",\"⊃\",\"ⓛ\",\"‚\",\"♰\",\"́\",\"✏\",\"⏑\",\"̶\",\"ⓢ\",\"⩾\",\"￠\",\"❍\",\"≃\",\"⋰\",\"♋\",\"､\",\"̂\",\"❋\",\"✳\",\"ⓤ\",\"╤\",\"▕\",\"⌣\",\"✸\",\"℮\",\"⁺\",\"▨\",\"╨\",\"ⓥ\",\"♈\",\"❃\",\"☝\",\"✻\",\"⊇\",\"≻\",\"♘\",\"♞\",\"◂\",\"✟\",\"⌠\",\"✠\",\"☚\",\"✥\",\"❊\",\"ⓒ\",\"⌈\",\"❅\",\"ⓡ\",\"♧\",\"ⓞ\",\"▭\",\"❱\",\"ⓣ\",\"∟\",\"☕\",\"♺\",\"∵\",\"⍝\",\"ⓑ\",\"✵\",\"✣\",\"٭\",\"♆\",\"ⓘ\",\"∶\",\"⚜\",\"◞\",\"்\",\"✹\",\"➥\",\"↕\",\"̳\",\"∷\",\"✋\",\"➧\",\"∋\",\"̿\",\"ͧ\",\"┅\",\"⥤\",\"⬆\",\"⋱\",\"☄\",\"↖\",\"⋮\",\"۔\",\"♌\",\"ⓛ\",\"╕\",\"♓\",\"❯\",\"♍\",\"▋\",\"✺\",\"⭐\",\"✾\",\"♊\",\"➣\",\"▿\",\"ⓑ\",\"♉\",\"⏠\",\"◾\",\"▹\",\"⩽\",\"↦\",\"╥\",\"⍵\",\"⌋\",\"։\",\"➨\",\"∮\",\"⇥\",\"ⓗ\",\"ⓓ\",\"⁻\",\"⎝\",\"⌥\",\"⌉\",\"◔\",\"◑\",\"✼\",\"♎\",\"♐\",\"╪\",\"⊚\",\"☒\",\"⇤\",\"ⓜ\",\"⎠\",\"◐\",\"⚠\",\"╞\",\"◗\",\"⎕\",\"ⓨ\",\"☟\",\"ⓟ\",\"♟\",\"❈\",\"↬\",\"ⓓ\",\"◻\",\"♮\",\"❙\",\"♤\",\"∉\",\"؛\",\"⁂\",\"ⓝ\",\"־\",\"♑\",\"╫\",\"╓\",\"╳\",\"⬅\",\"☔\",\"☸\",\"┄\",\"╧\",\"׃\",\"⎢\",\"❆\",\"⋄\",\"⚫\",\"̏\",\"☏\",\"➞\",\"͂\",\"␙\",\"ⓤ\",\"◟\",\"̊\",\"⚐\",\"✙\",\"↙\",\"̾\",\"℘\",\"✷\",\"⍺\",\"❌\",\"⊢\",\"▵\",\"✅\",\"ⓖ\",\"☨\",\"▰\",\"╡\",\"ⓜ\",\"☤\",\"∽\",\"╘\",\"˹\",\"↨\",\"♙\",\"⬇\",\"♱\",\"⌡\",\"⠀\",\"╛\",\"❕\",\"┉\",\"ⓟ\",\"̀\",\"♖\",\"ⓚ\",\"┆\",\"⎜\",\"◜\",\"⚾\",\"⤴\",\"✇\",\"╟\",\"⎛\",\"☩\",\"➲\",\"➟\",\"ⓥ\",\"ⓗ\",\"⏝\",\"◃\",\"╢\",\"↯\",\"✆\",\"˃\",\"⍴\",\"❇\",\"⚽\",\"╒\",\"̸\",\"♜\",\"☓\",\"➳\",\"⇄\",\"☬\",\"⚑\",\"✐\",\"⌃\",\"◅\",\"▢\",\"❐\",\"∊\",\"☈\",\"॥\",\"⎮\",\"▩\",\"ு\",\"⊹\",\"‵\",\"␔\",\"☊\",\"➸\",\"̌\",\"☿\",\"⇉\",\"⊳\",\"╙\",\"ⓦ\",\"⇣\",\"｛\",\"̄\",\"↝\",\"⎟\",\"▍\",\"❗\",\"״\",\"΄\",\"▞\",\"◁\",\"⛄\",\"⇝\",\"⎪\",\"♁\",\"⇠\",\"☇\",\"✊\",\"ி\",\"｝\",\"⭕\",\"➘\",\"⁀\",\"☙\",\"❛\",\"❓\",\"⟲\",\"⇀\",\"≲\",\"ⓕ\",\"⎥\",\"\\u06dd\",\"ͤ\",\"₋\",\"̱\",\"̎\",\"♝\",\"≳\",\"▙\",\"➭\",\"܀\",\"ⓖ\",\"⇛\",\"▊\",\"⇗\",\"̷\",\"⇱\",\"℅\",\"ⓧ\",\"⚛\",\"̐\",\"̕\",\"⇌\",\"␀\",\"≌\",\"ⓦ\",\"⊤\",\"̓\",\"☦\",\"ⓕ\",\"▜\",\"➙\",\"ⓨ\",\"⌨\",\"◮\",\"☷\",\"◍\",\"ⓚ\",\"≔\",\"⏩\",\"⍳\",\"℞\",\"┋\",\"˻\",\"▚\",\"≺\",\"ْ\",\"▟\",\"➻\",\"̪\",\"⏪\",\"̉\",\"⎞\",\"┇\",\"⍟\",\"⇪\",\"▎\",\"⇦\",\"␝\",\"⤷\",\"≖\",\"⟶\",\"♗\",\"̴\",\"♄\",\"ͨ\",\"̈\",\"❜\",\"̡\",\"▛\",\"✁\",\"➩\",\"ா\",\"˂\",\"↥\",\"⏎\",\"⎷\",\"̲\",\"➖\",\"↲\",\"⩵\",\"̗\",\"❢\",\"≎\",\"⚔\",\"⇇\",\"̑\",\"⊿\",\"̖\",\"☍\",\"➹\",\"⥊\",\"⁁\",\"✢\"];\n\ndef clean_punct(x):\n    for punct in puncts:\n        if punct in x:\n            x = x.replace(punct, f' {punct} ')\n    return x\n\ndef data_cleaning(x):\n    x = clean_tag(x)\n    x = clean_contractions(x)\n    x = clean_punct(x)\n    return x\n\ntrain_data['preprocessed_question_text'] = train_data['question_text'].progress_map(lambda x: data_cleaning(x))\ntest_data['preprocessed_question_text'] = test_data['question_text'].progress_map(lambda x: data_cleaning(x))","metadata":{"scrolled":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Define the max sentence length. The length should be longer than most sentences in the dataset, otherwise it will lose a lot of useful features.","metadata":{}},{"cell_type":"code","source":"from transformers import BertConfig, BertTokenizer, TFBertModel\n\npretrained_model_name = \"bert-base-uncased\"\n\ntokenizer = BertTokenizer.from_pretrained(pretrained_model_name, do_lower_case=True)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batch_size = 64\nmax_length = 100\n\nX_data = tokenizer(\n    train_data['preprocessed_question_text'].tolist(), \n    max_length=max_length, \n    padding='max_length',\n    truncation=True,\n    return_tensors='np'\n)\n\nX_test = tokenizer(\n    test_data['preprocessed_question_text'].tolist(), \n    max_length=max_length, \n    padding='max_length',\n    truncation=True,\n    return_tensors='np'\n)\n\nX_data = {\n    \"input_ids\": X_data[\"input_ids\"],\n    \"token_type_ids\": X_data[\"token_type_ids\"],\n    \"attention_mask\": X_data[\"attention_mask\"],\n}\n\nX_test = {\n    \"input_ids\": X_test[\"input_ids\"],\n    \"token_type_ids\": X_test[\"token_type_ids\"],\n    \"attention_mask\": X_test[\"attention_mask\"],\n}\n\ny_data = train_data['target'].to_numpy().reshape(-1,1)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def create_model():\n    output_bias = tf.keras.initializers.Constant(np.log([positive/negative]))\n\n    config = BertConfig.from_pretrained(pretrained_model_name) \n    config.output_hidden_states=True\n\n    transformers_model = TFBertModel.from_pretrained(pretrained_model_name, config=config)\n    transformers_model.bert.trainable = False\n\n    input_ids = tf.keras.layers.Input(\n        shape=(max_length,), \n        name='input_ids', \n        dtype='int32'\n    )\n    input_token = tf.keras.layers.Input(\n        shape=(max_length,), \n        name='token_type_ids', \n        dtype='int32'\n    )\n    input_attention = tf.keras.layers.Input(\n        shape=(max_length,), \n        name='attention_mask', \n        dtype='int32'\n    )\n\n    x = transformers_model(input_ids=input_ids, token_type_ids=input_token, attention_mask=input_attention)\n    x = tf.keras.layers.concatenate(tuple([x.hidden_states[i] for i in [0, -2, -1]]))\n    \n    lstm = tf.keras.layers.Bidirectional(tf.keras.layers.LSTM(64, return_sequences=True))(x)\n    gru = tf.keras.layers.Bidirectional(tf.keras.layers.GRU(64, return_sequences=True))(x)\n    x = tf.keras.layers.Concatenate()([lstm, gru])\n    x = tf.keras.layers.GlobalAveragePooling1D()(x)\n\n    outputs = tf.keras.layers.Dense(1, activation='sigmoid', bias_initializer=output_bias)(x)\n\n    model = tf.keras.Model(inputs=[input_ids, input_attention, input_token], outputs=outputs)\n    \n    return model","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def f1_smart(y_true, y_pred):\n    args = np.argsort(y_pred)\n    tp = y_true.sum()\n    fs = (tp - np.cumsum(y_true[args[:-1]])) / np.arange(y_true.shape[0] + tp - 1, tp, -1)\n    res_idx = np.argmax(fs)\n    return 2 * fs[res_idx], (y_pred[args[res_idx]] + y_pred[args[res_idx + 1]]) / 2","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import tensorflow as tf\n\nstrategy = None\n\ntry:\n    tpu = tf.distribute.cluster_resolver.TPUClusterResolver.connect()\n    strategy = tf.distribute.TPUStrategy(tpu)\n    print('Use TPU')\nexcept ValueError:\n    if len(tf.config.list_physical_devices('GPU')) > 0:\n        strategy = tf.distribute.MirroredStrategy()\n        print('Use GPU')\n    else:\n        strategy = tf.distribute.get_strategy()\n        print('Use CPU')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.metrics import f1_score\nfrom sklearn.model_selection import StratifiedKFold\nfrom IPython.display import Image\nfrom keras.utils import plot_model\n\nweight_for_0 = (1 / negative) * (total) / 2.0 \nweight_for_1 = (1 / positive) * (total) / 2.0\n\nclass_weight = {0: weight_for_0, 1: weight_for_1}\n\ncheckpoint = tf.keras.callbacks.ModelCheckpoint('best_model.h5', monitor='val_loss', save_weights_only=True, save_best_only=True, mode='min')\n\nreduce_lr = tf.keras.callbacks.ReduceLROnPlateau(monitor='val_loss', factor=0.3, patience=1)\n\nkfold = StratifiedKFold(n_splits=5, shuffle=True, random_state=seed)\n\nwith strategy.scope():\n    model = create_model()\n    model.compile(loss='binary_crossentropy', optimizer='adam')\n    model.summary()\n    \n    for index, (train_index, valid_index) in enumerate(kfold.split(np.zeros(len(y_data)), y_data)):\n        if index > 1:\n            break\n            \n        y_train, y_val = y_data[train_index], y_data[valid_index]\n        X_train = {\n            \"input_ids\": X_data[\"input_ids\"][train_index],\n            \"token_type_ids\": X_data[\"token_type_ids\"][train_index],\n            \"attention_mask\": X_data[\"attention_mask\"][train_index],\n        }\n        X_val = {\n            \"input_ids\": X_data[\"input_ids\"][valid_index],\n            \"token_type_ids\": X_data[\"token_type_ids\"][valid_index],\n            \"attention_mask\": X_data[\"attention_mask\"][valid_index],\n        }\n\n        history = model.fit(\n            X_train,\n            y_train,\n            epochs=5,\n            batch_size=batch_size,\n            validation_data=(X_val, y_val),\n            class_weight=class_weight,\n            callbacks=[reduce_lr, checkpoint]\n        )\n\n        y_pred = model.predict(X_val)\n        f1, threshold = f1_smart(y_val, np.squeeze(y_pred))\n        print('Optimal F1: {:.4f} at threshold: {:.4f}\\n'.format(f1, threshold))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Predict on the test dataset and write to the file named as submission.csv.","metadata":{}},{"cell_type":"code","source":"with strategy.scope():\n    best_model = create_model()\n    best_model.load_weights('best_model.h5')\n\n    y_pred = model.predict(X_val)\n    f1, threshold = f1_smart(y_val, np.squeeze(y_pred))\n    print('Optimal F1: {:.4f} at threshold: {:.4f}\\n'.format(f1, threshold))\n    \n    Y_test = (best_model.predict(X_test) > threshold).astype(\"int32\")\n\n    print('Write results to submission.csv')\n    submit_data = pd.DataFrame({'qid': test_data.qid, 'prediction': Y_test.reshape(-1)})\n    submit_data.to_csv('submission.csv', index=False)\n\n!head submission.csv","metadata":{"scrolled":true,"trusted":true},"execution_count":null,"outputs":[]}]}