{"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":"# Disaster tweets","metadata":{}},{"cell_type":"markdown","source":"* idea - pretrain as a language model and fine tune for classification\n* model:\n 1. lstm with two modes: predicting next token and classifying tweets\n 2. alternating initial state based on keyword\n 3. regularization: \n  * regular dropout - done\n  * tied weights - done\n  * spatial dropout - done\n  * dropconnect\n  * shared dropout mask\n 4. pooling\n  ","metadata":{}},{"cell_type":"markdown","source":"TODO:\n* add pooling\n* add more dropouts\n* set up proper experiments","metadata":{}},{"cell_type":"markdown","source":"### Imports","metadata":{}},{"cell_type":"code","source":"import numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# utils\nfrom operator import itemgetter\nfrom typing import Tuple, List, Dict\nfrom tqdm import tqdm\nimport random\nimport copy\n\n# plotting\nimport matplotlib.pyplot as plt\n\n# text processing\nimport regex as re\nimport string\nfrom collections import Counter\n\n# model\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nfrom torch.nn.utils.rnn import pad_sequence, pack_sequence, pad_packed_sequence\n\n# training\nfrom sklearn.model_selection import train_test_split\nfrom torch.nn.functional import binary_cross_entropy\nfrom torch.nn import BCELoss\nfrom torch.nn import CrossEntropyLoss\nimport torch.optim as optim\n\nrandom_state = 42\ntorch.manual_seed(random_state)\nnp.random.seed(random_state)\nrandom.seed(random_state)\n\nif torch.cuda.is_available():\n    device = torch.device('cuda')\n    print(f'Using GPU device: {device}')\nelse:\n    device = torch.device('cpu')\n    print(f'GPU is not available, using CPU device {device}')\n","metadata":{"execution":{"iopub.status.busy":"2022-07-18T18:25:54.452498Z","iopub.execute_input":"2022-07-18T18:25:54.453120Z","iopub.status.idle":"2022-07-18T18:25:54.466190Z","shell.execute_reply.started":"2022-07-18T18:25:54.453063Z","shell.execute_reply":"2022-07-18T18:25:54.465202Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Data loading","metadata":{}},{"cell_type":"code","source":"import os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-07-18T18:25:54.739249Z","iopub.execute_input":"2022-07-18T18:25:54.739768Z","iopub.status.idle":"2022-07-18T18:25:54.752390Z","shell.execute_reply.started":"2022-07-18T18:25:54.739727Z","shell.execute_reply":"2022-07-18T18:25:54.751385Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = pd.read_csv('/kaggle/input/nlp-getting-started/train.csv')\ntest_df = pd.read_csv('/kaggle/input/nlp-getting-started/test.csv')\ndata_df = pd.concat([train_df, test_df]).reset_index(drop=True)","metadata":{"execution":{"iopub.status.busy":"2022-07-18T18:25:55.048185Z","iopub.execute_input":"2022-07-18T18:25:55.048603Z","iopub.status.idle":"2022-07-18T18:25:55.099436Z","shell.execute_reply.started":"2022-07-18T18:25:55.048567Z","shell.execute_reply":"2022-07-18T18:25:55.098726Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data preprocessing","metadata":{}},{"cell_type":"markdown","source":"### keywords preprocessing","metadata":{"execution":{"iopub.status.busy":"2022-07-02T16:41:44.710446Z","iopub.execute_input":"2022-07-02T16:41:44.710835Z","iopub.status.idle":"2022-07-02T16:41:44.715462Z","shell.execute_reply.started":"2022-07-02T16:41:44.710803Z","shell.execute_reply":"2022-07-02T16:41:44.714369Z"}}},{"cell_type":"code","source":"keywords = data_df.keyword\nkeywords = keywords.apply(lambda t: '<empty>' if type(t) != str and np.isnan(t) else t)\nkeywords.value_counts()","metadata":{"execution":{"iopub.status.busy":"2022-07-18T18:25:55.338832Z","iopub.execute_input":"2022-07-18T18:25:55.339540Z","iopub.status.idle":"2022-07-18T18:25:55.355674Z","shell.execute_reply.started":"2022-07-18T18:25:55.339506Z","shell.execute_reply":"2022-07-18T18:25:55.354508Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Texts cleaning","metadata":{}},{"cell_type":"markdown","source":"TODO:\n* try lemmantization and stemming","metadata":{}},{"cell_type":"code","source":"all_texts = data_df.text\nall_texts_combined = '\\n'.join(data_df.text)\nsample_texts = all_texts.sample(10, random_state=random_state)\nsample_texts","metadata":{"execution":{"iopub.status.busy":"2022-07-18T18:25:55.621047Z","iopub.execute_input":"2022-07-18T18:25:55.621308Z","iopub.status.idle":"2022-07-18T18:25:55.632629Z","shell.execute_reply.started":"2022-07-18T18:25:55.621284Z","shell.execute_reply":"2022-07-18T18:25:55.631785Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class TextProcessor:\n    def __init__(self,\n                 remove_html=True, \n                 mask_links=True,\n                 remove_tags=False,\n                 mask_numbers=True,\n                 lower_case=True,\n                 indent_punctuation=True,\n                 remove_repeating_signs=True,\n                 mask_emoji=False,\n                 indent_dash=True,\n                 remove_quotations=True,\n                 ):\n        self.remove_html=remove_html\n        self.mask_links=mask_links\n        self.remove_tags=remove_tags\n        self.mask_numbers=mask_numbers\n        self.lower_case=lower_case\n        self.indent_punctuation=indent_punctuation\n        self.remove_repeating_signs=remove_repeating_signs\n        self.mask_emoji=mask_emoji\n        self.indent_dash=indent_dash\n        self.remove_quotations=remove_quotations\n\n    def _clean_text(self, text):\n        text = text.replace('%20', ' ')\n        if self.remove_html:\n            html = re.compile(r'<.*?>|&([a-z0-9]+|#[0-9]{1,6}|#x[0-9a-f]{1,6});')\n            text = re.sub(html, ' ', text)\n        if self.mask_links:\n            text = re.sub(r'https?://\\S+|www\\.\\S+', ' <link> ', text)\n        if self.remove_tags:\n            text = re.sub('#(?=\\w+)', ' ', text)\n            text = re.sub(r'@', ' ', text)\n        if self.remove_repeating_signs:\n            text = re.sub(r'[?]+', '?', text)\n            text = re.sub(r'[.]+', '.', text)\n            text = re.sub(r'!+', '!', text)\n        if self.mask_emoji:\n            emoji_pattern = re.compile(\n                '['\n                u'\\U0001F600-\\U0001F64F'  # emotions\n                u'\\U0001F300-\\U0001F5FF'  # symbols & pictographs\n                u'\\U0001F680-\\U0001F6FF'  # transport & map symbols\n                u'\\U0001F1E0-\\U0001F1FF'  # flags (iOS)\n                u'\\U00002702-\\U000027B0'\n                u'\\U000024C2-\\U0001F251'\n                ']+',\n                flags=re.UNICODE)\n            text = re.sub(emoji_pattern, '<emoji>', text)\n            text = re.sub(r':[(+)+D](?=[\\Z|\\s])', ' <emoji> ', text)\n        if self.indent_dash:\n            text = re.sub(r'-', ' - ', text)\n        if self.mask_numbers:\n            text = re.sub(r'[\\A|\\s]\\d+(\\S\\d+)?\\S{0,2}(?=[!|\\.|,|;|\\s+|\\Z])', ' <number> ', text)\n            text = re.sub(r'[\\^|\\s]\\d+[\\s|\\Z]', '<number>', text)\n        if self.remove_quotations:\n            text = re.sub(\"(\\s')|(^')\", r' ', text)\n            text = re.sub(r\"('\\s)|('\\Z)\", r' ', text)\n        if self.indent_punctuation:\n            text = re.sub(r'([@#.,!?():;])', r' \\1 ', text)\n        if self.mask_numbers:\n            text = re.sub(r'\\s\\d+\\s', ' <number> ', text)\n        if self.lower_case:\n            text = text.lower()\n        return text\n    \n    def __call__(self, text):\n        clean_text = self._clean_text(text)\n        return clean_text\n","metadata":{"execution":{"iopub.status.busy":"2022-07-18T18:25:55.974387Z","iopub.execute_input":"2022-07-18T18:25:55.974765Z","iopub.status.idle":"2022-07-18T18:25:55.990180Z","shell.execute_reply.started":"2022-07-18T18:25:55.974735Z","shell.execute_reply":"2022-07-18T18:25:55.989396Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"text_processor = TextProcessor(indent_dash=True, mask_emoji=False, mask_numbers=False)\nsample_texts_processed = sample_texts.apply(text_processor)\nprint(*[f'{idx}: {text[:144]}' for idx, text in enumerate(sample_texts)], sep='\\n')\nprint('-'*20)\nprint(*[f'{idx}: {text[:144]}' for idx, text in enumerate(sample_texts_processed)], sep='\\n')\n","metadata":{"execution":{"iopub.status.busy":"2022-07-18T18:25:56.393606Z","iopub.execute_input":"2022-07-18T18:25:56.394144Z","iopub.status.idle":"2022-07-18T18:25:56.404894Z","shell.execute_reply.started":"2022-07-18T18:25:56.394106Z","shell.execute_reply":"2022-07-18T18:25:56.403376Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"processed_texts = all_texts.apply(text_processor)","metadata":{"execution":{"iopub.status.busy":"2022-07-18T18:25:56.679967Z","iopub.execute_input":"2022-07-18T18:25:56.680291Z","iopub.status.idle":"2022-07-18T18:25:57.438684Z","shell.execute_reply.started":"2022-07-18T18:25:56.680266Z","shell.execute_reply":"2022-07-18T18:25:57.437895Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Tokenization","metadata":{}},{"cell_type":"code","source":"class Tokenizer:\n    def __init__(self, min_occurences=4,\n                 add_pad_token=True, add_eos_token=True):\n        self.min_occurences = min_occurences\n        self.token_to_id = None\n        self.id_to_token = None\n    \n    def build_vocab(self, texts):\n        if type(texts) is not str:\n            texts = ' '.join(texts)\n        # find all tokens\n        tokens = texts.split()\n        n_tokens = len(tokens)\n        token_cnts = Counter(tokens).most_common()\n        tokens = [x[0] for x in token_cnts if x[1] >= self.min_occurences]\n        tokens += ['<unk>']\n        tokens += ['<eos>']\n        # build dicts\n        self.token_to_id = {k: v for k,v in zip(tokens, range(1, len(tokens)+1))}\n        self.token_to_id['<pad>'] = 0\n        self.id_to_token = {k: v for v, k in self.token_to_id.items()}\n        return token_cnts, n_tokens\n        \n    def encode(self, text: str, add_eos=False):\n        if self.token_to_id is None:\n            raise Exception(\"Token dicts haven't been initialized yet\")\n        tokens = text.split()\n        if add_eos:\n            tokens += ['<eos>']\n        ids = [self.token_to_id[token] if token in self.token_to_id else self.token_to_id['<unk>']\n               for token in tokens]\n        return ids\n    \n    def decode(self, token_ids):\n        if self.token_to_id is None:\n            raise Exception(\"Token dicts haven't been initialized yet\")\n        return [self.id_to_token[token_id] for token_id in token_ids]\n    \n    \n    def __len__(self):\n        return len(self.token_to_id)\n    \n    \n    def __call__(self, text, add_eos=False):\n        return self.encode(text, add_eos)\n","metadata":{"execution":{"iopub.status.busy":"2022-07-18T18:25:57.440344Z","iopub.execute_input":"2022-07-18T18:25:57.440704Z","iopub.status.idle":"2022-07-18T18:25:57.452711Z","shell.execute_reply.started":"2022-07-18T18:25:57.440654Z","shell.execute_reply":"2022-07-18T18:25:57.451802Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"min_occurences=3\n\ntokenizer = Tokenizer(min_occurences=min_occurences)\ntoken_cnts, n_tokens = tokenizer.build_vocab(processed_texts)\n\nprint(f'tocken count: {n_tokens}')\nprint(f'unique tokens count: {len(token_cnts)}')\nprint(f'n words occuring 1 time: {len([k for k,v in token_cnts if v == 1])}')\nprint(f'n words occuring 2 times: {len([k for k,v in token_cnts if v == 2])}')\nprint(f'n words occuring 3 times: {len([k for k,v in token_cnts if v == 3])}')\nprint(f'n words occuring 4 times: {len([k for k,v in token_cnts if v == 4])}')\nprint(f'n words occuring 5 times: {len([k for k,v in token_cnts if v == 5])}')\nprint(f'n words occuring 6 or more times: {len([k for k,v in token_cnts if v >= 6])}')\n\n","metadata":{"execution":{"iopub.status.busy":"2022-07-18T18:25:57.454133Z","iopub.execute_input":"2022-07-18T18:25:57.454717Z","iopub.status.idle":"2022-07-18T18:25:57.530484Z","shell.execute_reply.started":"2022-07-18T18:25:57.454661Z","shell.execute_reply":"2022-07-18T18:25:57.529533Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# n most common tokens\nn_tokens = 30\nplt.figure(figsize=(8, 8))\nplt.barh(*zip(*sorted(token_cnts[:n_tokens], key=itemgetter(1))))\nplt.grid(visible=True, alpha=0.5)","metadata":{"execution":{"iopub.status.busy":"2022-07-18T18:25:57.826386Z","iopub.execute_input":"2022-07-18T18:25:57.826771Z","iopub.status.idle":"2022-07-18T18:25:58.115549Z","shell.execute_reply.started":"2022-07-18T18:25:57.826741Z","shell.execute_reply":"2022-07-18T18:25:58.114767Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tokenized_sample = sample_texts_processed.apply(tokenizer)\n\nprint(*[f'{idx}: {text[:144]}' for idx, text in enumerate(sample_texts_processed)], sep='\\n')\nprint('-'*20)\nprint(*[f'{idx}: {\" \".join(tokenizer.decode(text[:144]))}' for idx, text in enumerate(tokenized_sample)], sep='\\n')","metadata":{"execution":{"iopub.status.busy":"2022-07-18T18:25:58.136171Z","iopub.execute_input":"2022-07-18T18:25:58.136450Z","iopub.status.idle":"2022-07-18T18:25:58.143077Z","shell.execute_reply.started":"2022-07-18T18:25:58.136425Z","shell.execute_reply":"2022-07-18T18:25:58.142225Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tokenized_texts = processed_texts.apply(tokenizer)","metadata":{"execution":{"iopub.status.busy":"2022-07-18T18:25:58.452254Z","iopub.execute_input":"2022-07-18T18:25:58.452554Z","iopub.status.idle":"2022-07-18T18:25:58.643249Z","shell.execute_reply.started":"2022-07-18T18:25:58.452526Z","shell.execute_reply":"2022-07-18T18:25:58.642474Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"keyword_tokenizer = Tokenizer(min_occurences=0)\nkeyword_tokenizer.build_vocab(keywords)\n\ntokenized_keywords = keywords.apply(keyword_tokenizer).apply(itemgetter(0))","metadata":{"execution":{"iopub.status.busy":"2022-07-18T18:48:00.975218Z","iopub.execute_input":"2022-07-18T18:48:00.976179Z","iopub.status.idle":"2022-07-18T18:48:01.006129Z","shell.execute_reply.started":"2022-07-18T18:48:00.976139Z","shell.execute_reply":"2022-07-18T18:48:01.005287Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{"execution":{"iopub.status.busy":"2022-06-26T09:59:25.230738Z","iopub.execute_input":"2022-06-26T09:59:25.231113Z","iopub.status.idle":"2022-06-26T09:59:25.315657Z","shell.execute_reply.started":"2022-06-26T09:59:25.231087Z","shell.execute_reply":"2022-06-26T09:59:25.314816Z"}}},{"cell_type":"markdown","source":"TODO:\n* add dropout - done\n* add tied embeddings - done\n* add augmentation to training data - done\n* add word dropout\n* (add shared state dropout)\n* (add dropconnect)\n* change default embeddings with GloVe or something\n* split model into modules","metadata":{}},{"cell_type":"code","source":"X_4d = torch.tensor([[[[1,2],[1,2],[1,2]], \n                      [[2,3],[2,3],[2,3]], \n                      [[3,4],[3,4],[3,4]]],\n                     [[[4,5],[4,5],[4,5]], \n                      [[5,6],[5,6],[5,6]], \n                      [[6,7],[6,7],[6,7]]], \n                     [[[7,8],[7,8],[7,8]], \n                      [[8,9],[8,9],[8,9]], \n                      [[9,9],[9,9],[9,9]]]], \n                    dtype=float)\n\nX_3d = torch.tensor([[[1,1,1], [2,2,2], [3,3,3]],\n                     [[4,4,4], [5,5,5], [6,6,6]],\n                     [[7,7,7], [8,8,8], [9,9,9]]], dtype=float)\n\nX_2d = torch.tensor([[1, 2, 3], [4, 5, 6], [7, 8, 9]], dtype=float)\n","metadata":{"execution":{"iopub.status.busy":"2022-07-18T18:25:59.075511Z","iopub.execute_input":"2022-07-18T18:25:59.075811Z","iopub.status.idle":"2022-07-18T18:25:59.085129Z","shell.execute_reply.started":"2022-07-18T18:25:59.075783Z","shell.execute_reply":"2022-07-18T18:25:59.084113Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SpatialDropout(nn.Module):\n    def __init__(self,\n                 dropout_rate: float=0):\n        super(SpatialDropout, self).__init__()\n        self.dropout = nn.Dropout(dropout_rate)\n    \n    def forward(self,\n                embs):\n        return self.dropout(torch.ones(embs.shape[:-1],\n                                       device=device, requires_grad=False)).unsqueeze(-1)*embs\n\n    \nclass MeanPooling(nn.Module):\n    def __init__(self):\n        raise Exception('Not implemented')\n        \n\nclass KeyStateEncoder(nn.Module):\n    def __init__(self,\n                 state_size,\n                 key_vocab_size,\n                 num_layers,\n                 word_dropout: float=0,\n                 dropout_rate: float=0):\n        super(KeyStateEncoder, self).__init__()\n        self.emb_size = 2*state_size\n        self.vocab_size = key_vocab_size\n        self.encoder = nn.Embedding(self.vocab_size, self.emb_size)\n        self.dropout = nn.Dropout(dropout_rate)\n        self.num_layers = num_layers\n    \n    def forward(self, keys):\n        embs = self.dropout(self.encoder(keys))\n        key_h0, key_c0 = embs.chunk(2, dim=-1)\n        # add empty states for layers above first\n        if self.num_layers > 1:\n            key_h0 = torch.stack([key_h0] + [torch.zeros_like(key_h0, requires_grad=False)\n                                             for i in range(1, self.num_layers)])\n            key_c0 = torch.stack([key_c0] + [torch.zeros_like(key_c0, requires_grad=False)\n                                             for i in range(1, self.num_layers)])\n        return key_h0, key_c0\n\n\nclass TweetsModel(nn.Module):\n    def __init__(self,\n                 emb_dim: int,\n                 hidden_size: int,\n                 vocab_size: int,\n                 key_vocab_size: int,\n                 key_weight: float=0.8,\n                 num_layers: int=1,\n                 dropout_rate: float=0,\n                 word_dropout_rate: float=0,\n                 tie_weights: bool=False\n                 ):\n        \n        super(TweetsModel, self).__init__()\n        self.emb_dim = emb_dim\n        self.num_layers = num_layers\n        self.hidden_size = hidden_size\n        self.vocab_size = vocab_size\n        self.encoder = nn.Embedding(num_embeddings=vocab_size,\n                                    embedding_dim=emb_dim)\n        self.spatial_dropout = SpatialDropout(word_dropout_rate)\n\n        self.key_state_encoder = KeyStateEncoder(hidden_size, key_vocab_size,\n                                                 num_layers, dropout_rate)\n        self.h0 = torch.empty((num_layers, hidden_size), requires_grad=True)\n        self.c0 = torch.empty((num_layers, hidden_size), requires_grad=True)\n\n        self.lstm = nn.LSTM(input_size=emb_dim,\n                            hidden_size=hidden_size,\n                            num_layers=num_layers,\n                            dropout=dropout_rate)\n        self.decoder = nn.Linear(in_features=hidden_size,\n                                 out_features=vocab_size)\n        self.classifier = nn.Linear(in_features=hidden_size,\n                                    out_features=1)\n        self.dropout = nn.Dropout(dropout_rate)\n        \n        if tie_weights:\n            if emb_dim != hidden_size:\n                raise ValueError('When using the tied flag, emb_dim must be equal to hidden_size')\n            self.decoder.weight = self.encoder.weight\n        \n        self.init_custom_weights()\n        \n    def init_custom_weights(self):\n        hid_stdv = 1.0 / np.sqrt(self.hidden_size)\n        emb_stdv = 1.0 / np.sqrt(self.emb_dim)\n        for weight in [self.h0, self.c0]:\n            nn.init.uniform_(weight, -hid_stdv, hid_stdv)\n        return 0\n    \n    \n    # classifier\n    def forward(self, tokens: torch.Tensor, \n                last_tokens_idxs: torch.Tensor, keywords: torch.Tensor):\n        batch_size = len(tokens)                \n        tokens_stepwise = tokens.transpose(0, 1)        \n        embs = self.spatial_dropout(self.encoder(tokens_stepwise))\n        \n        key_h0, key_c0 = self.key_state_encoder(keywords)\n        h0, c0 = self.h0, self.c0\n        h0 = h0.unsqueeze(1).repeat(1, batch_size, 1).to(device)\n        c0 = c0.unsqueeze(1).repeat(1, batch_size, 1).to(device)\n        \n        h0 = h0 + key_h0\n        c0 = c0 + key_c0\n\n        lstm_out, (hn, cn) = self.lstm(embs, (h0, c0))\n        res = lstm_out.transpose(0, 1)\n        \n        # gather output states for <eos> tokens\n        last_tokens_idxs = last_tokens_idxs.unsqueeze(-1).repeat(1, self.hidden_size).unsqueeze(1)\n        last_arr = torch.gather(input=res, dim=1, index=last_tokens_idxs)\n        \n        last_arr = self.dropout(last_arr)\n        return torch.sigmoid(self.classifier(last_arr).squeeze())\n    \n    \n    # lang model        \n    def forward_lm(self, sequences: torch.Tensor, state: (torch.Tensor, torch.Tensor)):\n        sequences_stepwise = sequences.transpose(0, 1)\n        h0, c0 = state\n        embs = self.spatial_dropout(self.encoder(sequences_stepwise))\n        lstm_out, (hn, cn) = self.lstm(embs, (h0, c0))\n        return self.decoder(lstm_out), (hn, cn)\n    \n    def get_init_state(self, batch_size, device):\n        h0 = self.h0.detach().clone()\n        c0 = self.c0.detach().clone()\n        h0 = h0.unsqueeze(1).repeat(1, batch_size, 1).to(device)\n        c0 = c0.unsqueeze(1).repeat(1, batch_size, 1).to(device)\n        return h0, c0\n","metadata":{"execution":{"iopub.status.busy":"2022-07-18T19:08:08.960180Z","iopub.execute_input":"2022-07-18T19:08:08.960529Z","iopub.status.idle":"2022-07-18T19:08:08.991690Z","shell.execute_reply.started":"2022-07-18T19:08:08.960499Z","shell.execute_reply":"2022-07-18T19:08:08.990922Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Pretraining as Language Model","metadata":{}},{"cell_type":"markdown","source":"### batch generation","metadata":{}},{"cell_type":"markdown","source":"1. shuffle tweets\n2. add \\<eos> to each tweet and contatinate\n3. split text into batch_size number of parts\n4. form batches by taking fixed length sequences from each text part","metadata":{}},{"cell_type":"code","source":"def lm_train_batch_generator(texts, token_to_id, batch_size, batch_sequence_len):\n    \n    shuffled_texts = texts.sample(frac=1).reset_index(drop=True)\n    shuffled_texts = shuffled_texts.apply(lambda t: t + [token_to_id['<eos>']])\n    text_sequence = shuffled_texts.sum()\n    \n    adjusted_n_words = len(text_sequence) - 1\n    part_size = adjusted_n_words // batch_size\n    batch_count = part_size // batch_sequence_len\n\n    for batch_idx in range(batch_count):\n        X = [text_sequence[batch_idx*batch_sequence_len + part_idx*part_size:\n                           (batch_idx+1)*batch_sequence_len + part_idx*part_size]\n             for part_idx in range(batch_size)]\n        Y = [text_sequence[batch_idx*batch_sequence_len + part_idx*part_size + 1:\n                           (batch_idx+1)*batch_sequence_len + part_idx*part_size + 1]\n             for part_idx in range(batch_size)]\n        \n        X = torch.tensor(X)\n        Y = torch.tensor(Y)\n        yield X, Y","metadata":{"execution":{"iopub.status.busy":"2022-07-18T18:26:00.355901Z","iopub.execute_input":"2022-07-18T18:26:00.356410Z","iopub.status.idle":"2022-07-18T18:26:00.364326Z","shell.execute_reply.started":"2022-07-18T18:26:00.356376Z","shell.execute_reply":"2022-07-18T18:26:00.363452Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### one epoch","metadata":{}},{"cell_type":"code","source":"def run_lm_epoch(model, optimizer, texts, loss_fn,\n                 token_to_id,\n                 device,\n                 batch_size=128,\n                 batch_sequence_len=16):\n    batch_gen = lm_train_batch_generator(texts, token_to_id,\n                                         batch_size=batch_size,\n                                         batch_sequence_len=batch_sequence_len)\n    total_loss, total_examples = 0.0, 0\n    state = model.get_init_state(batch_size, device)\n    \n    for step, (X, Y) in enumerate(batch_gen):\n        X = X.to(device)\n        Y = Y.to(device)\n        \n        state = [x.detach() for x in state]\n        logits, state = model.forward_lm(X, state)\n        logits_seqwise = logits.transpose(0, 1)\n        \n        loss = loss_fn(logits_seqwise.reshape((-1, model.vocab_size)), Y.reshape(-1))\n        total_examples += loss.size(0)\n        total_loss += loss.sum().item()\n        loss = loss.mean()\n        \n        if model.training:\n            optimizer.zero_grad()\n            loss.backward()\n            optimizer.step()\n            \n    return np.exp(total_loss / total_examples)","metadata":{"execution":{"iopub.status.busy":"2022-07-18T18:26:00.741489Z","iopub.execute_input":"2022-07-18T18:26:00.742062Z","iopub.status.idle":"2022-07-18T18:26:00.750811Z","shell.execute_reply.started":"2022-07-18T18:26:00.742024Z","shell.execute_reply":"2022-07-18T18:26:00.750112Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### train_lm","metadata":{}},{"cell_type":"code","source":"def train_lm(model, config, device, train_data, tokenizer, dev_data=None, verbose=True):\n    run_log = pd.DataFrame(columns=['epoch', 'lr', 'train_perp', 'dev_perp'])\n    \n    if config['optimizer'] == 'adam': \n        optimizer = optim.Adam(model.parameters(), lr=lm_train_config['lr'])\n    loss = torch.nn.CrossEntropyLoss(reduction='none')\n    \n    for i in range(1, config['num_epochs']+1):\n        # lr\n        lr_decay = config['lr_decay'] ** max(i - config['max_lr_epochs'], 0.0)\n        decayed_lr = config['lr'] * lr_decay\n        for g in optimizer.param_groups:\n            g['lr'] = decayed_lr\n        # training\n        model.train()\n        train_perplexity = run_lm_epoch(model, optimizer, train_data,loss_fn=loss, \n                                        device=device,\n                                        batch_size=config['batch_size'],\n                                        batch_sequence_len=config['batch_sequence_len'],\n                                        token_to_id=tokenizer.token_to_id)\n        # evaluation\n        dev_perplexity = None\n        if dev_data is not None:\n            model.eval()\n            with torch.no_grad():\n                dev_perplexity = run_lm_epoch(model, optimizer, dev_data,loss_fn=loss, \n                                              device=device,\n                                              batch_size=config['batch_size'],\n                                              batch_sequence_len=config['batch_sequence_len'],\n                                              token_to_id=tokenizer.token_to_id)\n        # logging\n        curr_epoch = {'epoch': i, 'lr': decayed_lr,\n                      'train_perp': train_perplexity, 'dev_perp': dev_perplexity}\n        run_log = run_log.append(curr_epoch, ignore_index=True)\n        if verbose:\n            print('epoch: {0:2d}, lr: {1:6.4f} train perp: {2:9.4f}'.format(i, round(decayed_lr, 4), round(train_perplexity, 4)), end=' ')\n            if dev_perplexity is not None:\n                print('dev perp: {1:9.4f}'.format(round(dev_perplexity, 4)))\n            else:\n                print()\n    \n    return model, run_log   ","metadata":{"execution":{"iopub.status.busy":"2022-07-18T18:26:01.041210Z","iopub.execute_input":"2022-07-18T18:26:01.041575Z","iopub.status.idle":"2022-07-18T18:26:01.053121Z","shell.execute_reply.started":"2022-07-18T18:26:01.041548Z","shell.execute_reply":"2022-07-18T18:26:01.052379Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training for classification","metadata":{}},{"cell_type":"markdown","source":"### batch generator","metadata":{}},{"cell_type":"markdown","source":"* sort tweets by length\n* combine sorted tweets' indexes into batches\n* randomly sort batches (with preference to short ones? can't google permutation with distribution (may split ds into len brackets and shuffle within them)\n* return batches in batch generator and pad elements within batch to equal length\n* add augmentation:\n * randomly choose p sequences in batch\n * split them in half (skip those which length is less than 3 (skip batches like this altogether))\n * randomly choose one of those half","metadata":{"execution":{"iopub.status.busy":"2022-07-03T08:47:05.879765Z","iopub.execute_input":"2022-07-03T08:47:05.880225Z","iopub.status.idle":"2022-07-03T08:47:05.892135Z","shell.execute_reply.started":"2022-07-03T08:47:05.880189Z","shell.execute_reply":"2022-07-03T08:47:05.891126Z"}}},{"cell_type":"code","source":"def classifier_train_batch_generator(texts, keywords, targets, token_to_id: dict, batch_size: int,\n                                     augmentation=False):\n    # returns keywords, padded_texts, text_lens, targets\n    tweet_idxs = sorted(range(len(texts)), key=lambda i: len(texts[i]))\n\n    n_texts = len(texts)\n    n_batches = n_texts // batch_size\n    \n    batch_ids = np.random.permutation(range(n_batches))\n    \n    for batch_id in batch_ids:\n        curr_idxs = tweet_idxs[batch_id*batch_size:\n                               (batch_id+1)*batch_size]\n        curr_keywords = torch.tensor(keywords[curr_idxs].to_numpy(dtype=int))\n        curr_texts = texts[curr_idxs].reset_index(drop=True)\n        curr_targets = torch.tensor(targets[curr_idxs])\n        text_max_len = len(curr_texts.iloc[-1])\n        \n        if augmentation and text_max_len > 4:\n            text_half_len = text_max_len // 2\n            curr_texts = curr_texts.apply(lambda t: t if not random.choice([True, False]) \n                                                    else t[text_half_len:] if random.choice([True, False])\n                                                    else t[:text_half_len])\n            \n        curr_texts = curr_texts.apply(lambda t: t + [token_to_id['<eos>']])\n\n        curr_text_lens = curr_texts.apply(len)\n        text_max_len = len(curr_text_lens)\n        curr_texts = pad_sequence([torch.tensor(text) for text in curr_texts], batch_first=True)\n        last_elems_idxs = torch.tensor(curr_text_lens.to_numpy()-1)\n        \n        yield curr_keywords, curr_texts, last_elems_idxs, curr_targets\n        ","metadata":{"execution":{"iopub.status.busy":"2022-07-18T18:40:27.524125Z","iopub.execute_input":"2022-07-18T18:40:27.525059Z","iopub.status.idle":"2022-07-18T18:40:27.535051Z","shell.execute_reply.started":"2022-07-18T18:40:27.525024Z","shell.execute_reply":"2022-07-18T18:40:27.534338Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### one epoch","metadata":{"execution":{"iopub.status.busy":"2022-07-02T11:10:30.430453Z","iopub.execute_input":"2022-07-02T11:10:30.430824Z","iopub.status.idle":"2022-07-02T11:10:30.453491Z","shell.execute_reply.started":"2022-07-02T11:10:30.43074Z","shell.execute_reply":"2022-07-02T11:10:30.452936Z"}}},{"cell_type":"code","source":"def run_cls_epoch(model, optimizer, texts, keywords, targets, loss_fn,\n                  token_to_id,\n                  device,\n                  batch_size=128, augmentation=True):\n                \n    batch_gen = classifier_train_batch_generator(texts=texts, keywords=keywords, targets=targets,\n                                                 token_to_id=token_to_id,\n                                                 batch_size=batch_size,\n                                                 augmentation=augmentation)\n    epoch_loss = 0\n    epoch_accuracy = 0\n    step = 0\n    all_preds = []\n    \n    for keywords, texts, last_idxs, targets in batch_gen:\n        \n        texts = texts.to(device)\n        last_idxs = last_idxs.to(device)\n        keywords = keywords.to(device)\n        targets = targets.to(device)\n        \n        preds = model(texts, last_idxs, keywords)\n        loss = loss_fn(preds, targets)\n        epoch_loss += loss.item()\n\n        if model.training:\n            optimizer.zero_grad()\n            loss.backward()\n            optimizer.step()\n        step += 1\n        \n        epoch_accuracy += sum(np.round_(np.array(preds.detach().cpu())) \\\n                              == np.array(targets.cpu()))/len(preds)\n    \n    epoch_ce = epoch_loss / step\n    epoch_acc = epoch_accuracy / step\n    return epoch_ce, epoch_acc","metadata":{"execution":{"iopub.status.busy":"2022-07-18T18:33:50.887531Z","iopub.execute_input":"2022-07-18T18:33:50.888214Z","iopub.status.idle":"2022-07-18T18:33:50.898312Z","shell.execute_reply.started":"2022-07-18T18:33:50.888117Z","shell.execute_reply":"2022-07-18T18:33:50.897321Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### cls_train","metadata":{}},{"cell_type":"code","source":"# training\ndef train_cls(model, config, device, tokenizer,\n              train_texts, train_keywords, train_targets,\n              dev_texts=None, dev_keywords=None, dev_targets=None,\n              verbose=True):\n    run_log = pd.DataFrame(columns=['epoch', 'lr', 'train_ce', 'train_acc',\n                                    'dev_ce', 'dev_acc'])\n    if config['optimizer'] == 'adam':\n        optimizer = optim.Adam(model.parameters(), lr=config['lr'])\n    loss = BCELoss(weight=None, reduction='mean')\n\n    for i in range(1, config['num_epochs']+1):\n        #lr\n        lr_decay = config['lr_decay'] ** max(i - config['max_lr_epochs'], 0.0)\n        decayed_lr = config['lr'] * lr_decay\n        for g in optimizer.param_groups:\n            g['lr'] = decayed_lr\n        # train\n        model.train()\n        train_ce, train_acc = run_cls_epoch(model, optimizer, train_texts, train_keywords, train_targets,\n                                            loss_fn=loss, batch_size=config['batch_size'],\n                                            token_to_id=tokenizer.token_to_id, augmentation=True, device=device)\n        # evaluation\n        dev_ce, dev_acc = None, None\n        if all([x is not None for x in [dev_texts, dev_keywords, dev_targets]]):\n            model.eval()\n            with torch.no_grad():\n                dev_ce, dev_acc = run_cls_epoch(model, optimizer, dev_texts, dev_keywords, dev_targets,\n                                                loss_fn=loss, batch_size=config['batch_size'],\n                                                token_to_id=tokenizer.token_to_id, augmentation=False, device=device)\n        # logging\n        curr_epoch = {'epoch': i, 'lr': decayed_lr,\n                      'train_ce': train_ce, 'train_acc': train_acc,\n                      'dev_ce': dev_ce, 'dev_acc': dev_acc}\n        run_log = run_log.append(curr_epoch, ignore_index=True)\n        if verbose:\n            if i % 2: continue\n            print(f'epoch: {i}, lr: {round(decayed_lr, 5)}',\n                  f'train acc: {round(train_acc, 4)} dev acc: {round(dev_acc, 4) if dev_acc is not None else None}',\n                  f'train ce: {round(train_ce, 4)}, dev ce: {round(dev_ce, 4) if dev_ce is not None else None}',\n                  '-'*10, sep='\\n')\n        \n    return model, run_log","metadata":{"execution":{"iopub.status.busy":"2022-07-18T18:33:52.966657Z","iopub.execute_input":"2022-07-18T18:33:52.967470Z","iopub.status.idle":"2022-07-18T18:33:52.980608Z","shell.execute_reply.started":"2022-07-18T18:33:52.967436Z","shell.execute_reply":"2022-07-18T18:33:52.979387Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Experiments and hyperparameter tuning","metadata":{}},{"cell_type":"markdown","source":"### language model","metadata":{}},{"cell_type":"markdown","source":"TODO: \n* train on increasing sequence lengths\n* make training reproducible each time","metadata":{}},{"cell_type":"code","source":"model_config = {'num_layers': 3,\n                'emb_size': 192, 'hidden_size': 192,\n                'vocab_size': len(tokenizer),\n                'keyword_vocab_size': len(keyword_tokenizer)+1,\n                'dropout_rate': 0.4,\n                'word_dropout_rate': 0.3,\n                'tie_weights': True}\n\nmodel = TweetsModel(emb_dim=model_config['emb_size'], hidden_size=model_config['hidden_size'],\n                    vocab_size=model_config['vocab_size'], key_vocab_size=model_config['keyword_vocab_size'],\n                    num_layers=model_config['num_layers'], dropout_rate=model_config['dropout_rate'],\n                    word_dropout_rate=model_config['word_dropout_rate'],\n                    tie_weights=model_config['tie_weights'])\nmodel.to(device);","metadata":{"execution":{"iopub.status.busy":"2022-07-18T18:56:07.448389Z","iopub.execute_input":"2022-07-18T18:56:07.448750Z","iopub.status.idle":"2022-07-18T18:56:07.487700Z","shell.execute_reply.started":"2022-07-18T18:56:07.448718Z","shell.execute_reply":"2022-07-18T18:56:07.486969Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# data\ntrain_texts, dev_texts = train_test_split(tokenized_texts, test_size=0.3)\n\n# training\nlm_train_config = {'lr': 0.003,\n                    'batch_size': 512,\n                    'batch_sequence_len': 16,\n                    'lr_decay': 0.95,\n                    'max_lr_epochs': 8,\n                    'num_epochs': 16,\n                    'optimizer': 'adam'}\n\nmodel, train_lm_log = train_lm(model, lm_train_config, device, train_texts, tokenizer, dev_texts, verbose=False)\n\nplt.figure(figsize=(12, 8))\nplt.grid(visible=True)\nplt.plot(train_lm_log['epoch'], train_lm_log['train_perp'])\nplt.plot(train_lm_log['epoch'], train_lm_log['dev_perp'])\nplt.legend(['train_perp', 'dev_perp']);","metadata":{"execution":{"iopub.status.busy":"2022-07-18T18:56:09.239896Z","iopub.execute_input":"2022-07-18T18:56:09.240874Z","iopub.status.idle":"2022-07-18T18:56:42.512395Z","shell.execute_reply.started":"2022-07-18T18:56:09.240802Z","shell.execute_reply":"2022-07-18T18:56:42.511673Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# copy pretrained model\nlm_trained_model = copy.deepcopy(model)","metadata":{"execution":{"iopub.status.busy":"2022-07-18T18:56:42.514059Z","iopub.execute_input":"2022-07-18T18:56:42.514617Z","iopub.status.idle":"2022-07-18T18:56:42.521547Z","shell.execute_reply.started":"2022-07-18T18:56:42.514569Z","shell.execute_reply":"2022-07-18T18:56:42.520817Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### classifier","metadata":{}},{"cell_type":"code","source":"# model\nuse_pretrained_model = True\n\n\nif use_pretrained_model:\n    model = copy.deepcopy(lm_trained_model)\nelse:\n    model_config = {'num_layers': 4,\n                    'emb_size': 128, 'hidden_size': 128,\n                    'vocab_size': len(tokenizer), \n                    'keyword_vocab_size': len(keyword_tokenizer)+1,\n                    'dropout_rate': 0.4,\n                    'word_dropout_rate': 0.3,\n                    'tie_weights': True}\n    model = TweetsModel(emb_dim=model_config['emb_size'], hidden_size=model_config['hidden_size'],\n                        vocab_size=model_config['vocab_size'], key_vocab_size=model_config['keyword_vocab_size'],\n                        num_layers=model_config['num_layers'], dropout_rate=model_config['dropout_rate'],\n                        word_dropout_rate=model_config['word_dropout_rate'],\n                        tie_weights=model_config['tie_weights'])\nmodel.to(device);","metadata":{"execution":{"iopub.status.busy":"2022-07-18T18:56:43.974260Z","iopub.execute_input":"2022-07-18T18:56:43.974630Z","iopub.status.idle":"2022-07-18T18:56:43.987617Z","shell.execute_reply.started":"2022-07-18T18:56:43.974600Z","shell.execute_reply":"2022-07-18T18:56:43.986809Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# data\ntrain_idxs = list(range(len(train_df)))\ny = train_df.target.to_numpy(dtype=np.float32)\n\nX_train_idxs, X_dev_idxs, y_train, y_dev = train_test_split(train_idxs, y, test_size=0.3)\n\ntrain_texts = tokenized_texts[X_train_idxs].reset_index(drop=True)\ntrain_keywords = tokenized_keywords[X_train_idxs].reset_index(drop=True)\ndev_texts = tokenized_texts[X_dev_idxs].reset_index(drop=True)\ndev_keywords = tokenized_keywords[X_dev_idxs].reset_index(drop=True)\n\n\n# training\ncls_train_config = {'lr': 0.0005,\n                    'batch_size': 512,\n                    'lr_decay': 0.95,\n                    'max_lr_epochs': 10,\n                    'num_epochs': 50,\n                    'optimizer': 'adam'}\n\nmodel, train_cls_log = train_cls(model=model, config=cls_train_config, device=device, tokenizer=tokenizer,\n                                 train_texts=train_texts, train_keywords=train_keywords, train_targets=y_train,\n                                 dev_texts=dev_texts, dev_keywords=dev_keywords, dev_targets=y_dev,\n                                 verbose=False)\n\nf, ax = plt.subplots(1, 2, figsize=(16, 6))\nax[0].plot(train_cls_log['epoch'], train_cls_log['train_ce'])\nax[0].plot(train_cls_log['epoch'], train_cls_log['dev_ce'])\nax[0].legend(['train_ce', 'dev_ce'])\n\nax[1].plot(train_cls_log['epoch'], train_cls_log['train_acc'])\nax[1].plot(train_cls_log['epoch'], train_cls_log['dev_acc'])\nax[1].legend(['train_acc', 'dev_acc']);","metadata":{"execution":{"iopub.status.busy":"2022-07-18T18:56:45.563412Z","iopub.execute_input":"2022-07-18T18:56:45.563764Z","iopub.status.idle":"2022-07-18T18:57:04.323580Z","shell.execute_reply.started":"2022-07-18T18:56:45.563733Z","shell.execute_reply":"2022-07-18T18:57:04.322830Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Submission","metadata":{}},{"cell_type":"markdown","source":"#### config","metadata":{}},{"cell_type":"code","source":"model_config = {'num_layers': 3,\n                'emb_size': 192, 'hidden_size': 192,\n                'vocab_size': len(tokenizer),\n                'keyword_vocab_size': len(keyword_tokenizer)+1,\n                'dropout_rate': 0.4,\n                'word_dropout_rate': 0.3,\n                'tie_weights': True}\n\nlm_train_config = {'lr': 0.003,\n                    'batch_size': 512,\n                    'batch_sequence_len': 16,\n                    'lr_decay': 0.9,\n                    'max_lr_epochs': 8,\n                    'num_epochs': 16,\n                    'optimizer': 'adam'}\n\ncls_train_config = {'lr': 0.0005,\n                    'batch_size': 512,\n                    'lr_decay': 0.95,\n                    'max_lr_epochs': 10,\n                    'num_epochs': 12,\n                    'optimizer': 'adam'}\n","metadata":{"execution":{"iopub.status.busy":"2022-07-18T18:27:08.759768Z","iopub.status.idle":"2022-07-18T18:27:08.760373Z","shell.execute_reply.started":"2022-07-18T18:27:08.760137Z","shell.execute_reply":"2022-07-18T18:27:08.760162Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### init model","metadata":{}},{"cell_type":"code","source":"model = TweetsModel(emb_dim=model_config['emb_size'], hidden_size=model_config['hidden_size'],\n                    vocab_size=model_config['vocab_size'], key_vocab_size=model_config['keyword_vocab_size'],\n                    num_layers=model_config['num_layers'], dropout_rate=model_config['dropout_rate'],\n                    word_dropout_rate=model_config['word_dropout_rate'],\n                    tie_weights=model_config['tie_weights'])\n\nmodel.to(device);","metadata":{"execution":{"iopub.status.busy":"2022-07-18T18:27:08.761470Z","iopub.status.idle":"2022-07-18T18:27:08.762081Z","shell.execute_reply.started":"2022-07-18T18:27:08.761828Z","shell.execute_reply":"2022-07-18T18:27:08.761865Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### pretrain","metadata":{}},{"cell_type":"code","source":"train_texts = tokenized_texts\n\nmodel, lm_run_log = train_lm(model, lm_train_config, device, train_texts, tokenizer)","metadata":{"execution":{"iopub.status.busy":"2022-07-18T18:27:08.763144Z","iopub.status.idle":"2022-07-18T18:27:08.763727Z","shell.execute_reply.started":"2022-07-18T18:27:08.763486Z","shell.execute_reply":"2022-07-18T18:27:08.763510Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### finetune","metadata":{}},{"cell_type":"code","source":"train_texts = tokenized_texts[:len(train_df)].reset_index(drop=True)\ntrain_keywords = tokenized_keywords[:len(train_df)].reset_index(drop=True)\ntrain_targets = train_df.target.to_numpy(dtype=np.float32)\n\nmodel, cls_run_logs = train_cls(model, cls_train_config, device, tokenizer,\n                                train_texts, train_keywords, train_targets)","metadata":{"execution":{"iopub.status.busy":"2022-07-18T18:27:08.764790Z","iopub.status.idle":"2022-07-18T18:27:08.765380Z","shell.execute_reply.started":"2022-07-18T18:27:08.765152Z","shell.execute_reply":"2022-07-18T18:27:08.765175Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### predictions","metadata":{}},{"cell_type":"code","source":"# adequacy test\nmodel.eval()\nwith torch.no_grad():\n    preds = [model(torch.tensor(text).unsqueeze(0).to(device),\n                   torch.tensor(len(text)-1).unsqueeze(0).to(device),\n                   torch.tensor(keyword).unsqueeze(0).to(device)).item()\n             for (text, keyword) in zip(train_texts, train_keywords)]\n\nsum(np.round_(np.array(preds)) == y) / len(preds)","metadata":{"execution":{"iopub.status.busy":"2022-07-18T18:27:08.766438Z","iopub.status.idle":"2022-07-18T18:27:08.767029Z","shell.execute_reply.started":"2022-07-18T18:27:08.766787Z","shell.execute_reply":"2022-07-18T18:27:08.766810Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_texts = tokenized_texts[len(train_df):]\ntest_keywords = tokenized_keywords[len(train_df):]\n\nmodel.eval()\nwith torch.no_grad():\n    preds = [model(torch.tensor(text).unsqueeze(0).to(device),\n                   torch.tensor(len(text)-1).unsqueeze(0).to(device),\n                   torch.tensor(keyword).unsqueeze(0).to(device)).item()\n             for (text, keyword) in zip(test_texts, test_keywords)]\n\npreds = np.round_(np.array(preds)).astype(int)\noutput = pd.DataFrame({'id': test_df['id'], 'target': preds})\noutput.to_csv(\"submission.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2022-07-18T18:27:08.768089Z","iopub.status.idle":"2022-07-18T18:27:08.768659Z","shell.execute_reply.started":"2022-07-18T18:27:08.768432Z","shell.execute_reply":"2022-07-18T18:27:08.768455Z"},"trusted":true},"execution_count":null,"outputs":[]}]}