{"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":"## Experiment log:\n- First, we took only the tweet text, without keywords or locations, and tried a simple model with an embedding layer, followed by a mean layer, then the output layers. Stopwords were not removed and no stemming was performed. Accuracy was varying between runs, ranging between 0.625 0.96875. We can do better.\n- Next,we tried removing stopwords. Accuracy did not peak as high as it was before, but the variance was less. Acuracy ranged between 0.46875 and 0.8125.\n- We tried stemming alone. Same as removing stopwords. Acuracy ranged between 0.5 and 0.8125.\n- Both stemming and removing stopwords made the training smoother, with the accuracy less fluctuating and almost steadily increasing. Range was between 0.46875 and 0.75.\n- After looking at kewyords, they seem helpful. They are the one or few words that are the main focus of the tweet. We could use that.\n- We tried just appending the keywords at the end of the original tweets, althought they are there anyway. We thought that would help the model pay more attention to the keyword and that would help capturing the sentiment of the tweet. It could have done that, but we did not see any significant increase in accuracy, it was between 0.5 and 0.8125. We need to think of another way to include them.\n-  Next, we tried an Embedding, LSTM, Fn (select_last), Dense (output), LogSoftmax/Sigmoid/Relu architecture. It was terrible! The accuracy was 0.5 most of the time, and even dropped to 0.40625 briefly. \n- We tried then to be creative :D we created two branches of Embedding, Mean, Dense, and LogSoftmax, one to process the tweet text and the other to process the keyword alone, then we averaged the log-softmax scores. The results were much better! The accuracy peaked at 1.0 briefly. We then did some tweaking of the learning rate and number of iterations not just to get higher accuracy but also to make it more steady and consistent.\n- Locations seemed interesting, I tried looking up location values in a third party repository for geolocation data. We were particularly interested in longitude and latitude. I got that, converted long/lat to spherical coordinates, fed that to a third branch in the model, and averaged the scores. The results were terrible! The model accuracy struggled to stay between mid 50s and mid 40s. It even dropped below that range. I tried adding another hidden layer between the input and output layers in that branch but that did not help. The problem I could see is that there were multiple suggestions for so many locations. I tried choosing one at random and choosing the first one, but neither helped. Also, imputing missing location values with 0s for lon and lat could have made things worse, because that's an actual geographical location that the model tried to relate to disasters.\n- I tried reading location like I read text and keyword and using the vocabulary to transform it into a a tensor, and run it through yet another branch. Results were better that the geolcation data not really impressive.\n- Next thing to try was to concatenate keywords and locations to the text with separators in between, and use the original architecture with Embedding, Mean, Dense, and LogSoftmax layers. It worked. Average accuracy was 0.691\n- I really wanted to make this better, I added key word to location and fed it to a separate branch, and the tweet body to the other branch, we could squeeze a few more hundredths of degrees of accuracy and F1 scores, average accuracy was   \n- Eureka! Transfer learning. I downloaded Glove pre-trained word-embedding, and intilaized the embedding layer with those weights and shared it between the two branches. Because the layer is trainable, the weights were overwritten during the training process. We need to lock it down and make it non-trainable. The problem is that Trax does not have such layer, we will have to rig something up.\n- I tried to overload backward() and has_backward in a custom Embedding layer to make it untrainable but it does not seem to work. The best we could achive so far was training an embedding layer in two branches one for the tweet body and the other for the keyword + location\n- I went back to using the geolocation data. I transformed long/lat data to 3d spherical coordinates and fed it to the model in a separate branch. The score was not much different, but the model is more convencing to me.","metadata":{}},{"cell_type":"code","source":"!pip install trax","metadata":{"execution":{"iopub.status.busy":"2022-07-24T16:43:49.345814Z","iopub.execute_input":"2022-07-24T16:43:49.347022Z","iopub.status.idle":"2022-07-24T16:45:13.223123Z","shell.execute_reply.started":"2022-07-24T16:43:49.346902Z","shell.execute_reply":"2022-07-24T16:45:13.221057Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport re\nimport json\nimport math\nimport shutil\nimport string\nimport random\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport nltk\nfrom nltk.tokenize import TweetTokenizer\nfrom nltk.corpus import stopwords\nfrom nltk.stem import PorterStemmer\nimport trax\nimport trax.fastmath.numpy as np\nfrom trax import layers as tl\nfrom trax import optimizers\nfrom trax.supervised import training\n\n%matplotlib inline","metadata":{"execution":{"iopub.status.busy":"2022-07-24T16:46:11.277603Z","iopub.execute_input":"2022-07-24T16:46:11.278184Z","iopub.status.idle":"2022-07-24T16:47:24.636705Z","shell.execute_reply.started":"2022-07-24T16:46:11.278142Z","shell.execute_reply":"2022-07-24T16:47:24.635656Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import gzip\n# import pickle\n# import enchant","metadata":{"execution":{"iopub.status.busy":"2022-07-24T16:47:24.638684Z","iopub.execute_input":"2022-07-24T16:47:24.639835Z","iopub.status.idle":"2022-07-24T16:47:24.644273Z","shell.execute_reply.started":"2022-07-24T16:47:24.639800Z","shell.execute_reply":"2022-07-24T16:47:24.642977Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Loading tweets","metadata":{}},{"cell_type":"code","source":"MODEL_DIR = './model'\nOUTPUT_DIR = './'\n\nVAL_PCT = 0.2\nstopwords_english = stopwords.words('english')\n# dict_english = enchant.Dict('en_US')","metadata":{"execution":{"iopub.status.busy":"2022-07-24T17:01:42.062214Z","iopub.execute_input":"2022-07-24T17:01:42.062709Z","iopub.status.idle":"2022-07-24T17:01:42.069968Z","shell.execute_reply.started":"2022-07-24T17:01:42.062671Z","shell.execute_reply":"2022-07-24T17:01:42.068772Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# This part is relevant to the experiment with pretrained word embeddings.\n# I downloaded the Glove and tried to load them into a makeshasft non-trainable\n# Trax Embedding layer.\n# RAW_EMBED_FILE = '../../data/glove.twitter.27B.200d.txt'\n# PROC_EMBED_FILE = f'{OUTPUT_DIR}/embeddings.pkl.gz'\n\n\ndef parse_n_save_embeds(raw_embed_file, proc_embed_file):\n    embeds = {}\n    with open(raw_embed_file) as emf:\n        for line in emf:\n            tokens = line.split()\n            word = tokens[0].lower()\n            embed = [float(token) for token in tokens[1:]]\n            if dict_english.check(word):\n                if len(embed) == 200:\n                    embeds[word] = embed\n                else:\n                    print(f'>>> Word {word} has {len(embed)} size embedding! Ignoring.')\n                    \n    embed_matrix = np.array(list(embeds.values()), 'float32')\n        \n    word_idx = {w:i for i, w in enumerate(embeds.keys())}\n    word_idx['__unk__'] = len(word_idx)\n    word_idx['__pad__'] = len(word_idx)\n    \n    output = {\n        'flat_weights': [embed_matrix],\n        'flat_state': [], # Required later when initializing layer\n        'word_index': word_idx,\n    }\n    with gzip.open(proc_embed_file, 'wb') as ef:\n        pickle.dump(output, ef, protocol=pickle.HIGHEST_PROTOCOL)\n        \n    return embed_matrix, word_idx\n        \n\ndef load_embeds(proc_embed_file):\n    with gzip.open(proc_embed_file, 'rb') as ef:\n        embeds = pickle.load(ef)\n\n    embed_matrix = embeds['flat_weights']\n    word_index = embeds['word_index']\n    \n    return embed_matrix, word_index\n\n# EMBED_MATRIX, WORD_IDX = parse_n_save_embeds(RAW_EMBED_FILE, PROC_EMBED_FILE)\n# EMBED_MATRIX, WORD_IDX = load_embeds(PROC_EMBED_FILE)\n\n# EMBED_SIZE = len(EMBED_MATRIX[0][0])\n# print(f'>>> {len(EMBED_MATRIX[0])} word embeddings, each of size {EMBED_SIZE}')\n\n# IDX_WORD = {i:w for w, i in WORD_IDX.items()}\n\n# del EMBED_MATRIX","metadata":{"execution":{"iopub.status.busy":"2022-07-24T16:48:03.560693Z","iopub.execute_input":"2022-07-24T16:48:03.561235Z","iopub.status.idle":"2022-07-24T16:48:03.579934Z","shell.execute_reply.started":"2022-07-24T16:48:03.561194Z","shell.execute_reply":"2022-07-24T16:48:03.578938Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Loading the training and test sets\nall_train_tweets = pd.read_csv('../input/nlp-getting-started/train.csv', index_col=['id'])\nall_test_tweets = pd.read_csv('../input/nlp-getting-started/test.csv', index_col=['id'])","metadata":{"execution":{"iopub.status.busy":"2022-07-24T16:49:14.089560Z","iopub.execute_input":"2022-07-24T16:49:14.090095Z","iopub.status.idle":"2022-07-24T16:49:14.188653Z","shell.execute_reply.started":"2022-07-24T16:49:14.090058Z","shell.execute_reply":"2022-07-24T16:49:14.187578Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# This part is relevant to the experiment with geolcation data.\n# Tried to translate a location to long-lat, then transform it\n# to sperical coordinates\n# Check https://datascience.stackexchange.com/questions/13567/ways-to-deal-with-longitude-latitude-feature\ndef process_location(loc):\n    x = y = z = 0\n    data = loc.get('data')\n    if data and any([True if loc_data else False for loc_data in data]):\n        lat = float('inf')\n        lon = float('inf')\n        for loc_data in data:\n            loc_lat = loc_data['latitude']\n            loc_lon = loc_data['longitude']\n            if math.fabs(loc_lat) < math.fabs(lat):\n                lat = loc_lat\n                lon = loc_lon\n        lat = lat * math.pi / 180\n        lon = lon * math.pi / 180\n        x = math.cos(lat) * math.cos(lon)\n        y = math.cos(lat) * math.sin(lon) \n        z = math.sin(lat)\n    return [x, y, z]","metadata":{"execution":{"iopub.status.busy":"2022-07-24T16:49:17.263389Z","iopub.execute_input":"2022-07-24T16:49:17.264064Z","iopub.status.idle":"2022-07-24T16:49:17.273299Z","shell.execute_reply.started":"2022-07-24T16:49:17.264017Z","shell.execute_reply":"2022-07-24T16:49:17.272382Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tweet_geocode = pd.read_csv('../input/tweet-geocode/tweet_geocode.csv', index_col=['id'])\ntweet_geocode.location_clean = tweet_geocode.location_clean.fillna('{\"data\": []}').map(json.loads)\ntweet_geocode.location_clean = tweet_geocode.location_clean.apply(process_location)","metadata":{"execution":{"iopub.status.busy":"2022-07-24T16:50:33.569874Z","iopub.execute_input":"2022-07-24T16:50:33.570412Z","iopub.status.idle":"2022-07-24T16:50:34.159095Z","shell.execute_reply.started":"2022-07-24T16:50:33.570379Z","shell.execute_reply":"2022-07-24T16:50:34.158041Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"all_train_tweets.loc[all_train_tweets.target == 1].head()","metadata":{"execution":{"iopub.status.busy":"2022-07-24T16:50:35.375649Z","iopub.execute_input":"2022-07-24T16:50:35.376034Z","iopub.status.idle":"2022-07-24T16:50:35.399535Z","shell.execute_reply.started":"2022-07-24T16:50:35.376006Z","shell.execute_reply":"2022-07-24T16:50:35.398291Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"all_train_tweets.loc[all_train_tweets.target == 0].head()","metadata":{"execution":{"iopub.status.busy":"2022-07-24T16:50:37.512751Z","iopub.execute_input":"2022-07-24T16:50:37.513164Z","iopub.status.idle":"2022-07-24T16:50:37.526807Z","shell.execute_reply.started":"2022-07-24T16:50:37.513129Z","shell.execute_reply":"2022-07-24T16:50:37.525429Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"all_train_tweets.location.loc[~all_train_tweets.location.isna()].head()","metadata":{"execution":{"iopub.status.busy":"2022-07-24T16:50:39.182581Z","iopub.execute_input":"2022-07-24T16:50:39.183583Z","iopub.status.idle":"2022-07-24T16:50:39.195302Z","shell.execute_reply.started":"2022-07-24T16:50:39.183539Z","shell.execute_reply":"2022-07-24T16:50:39.194105Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"all_train_tweets.keyword.loc[~all_train_tweets.keyword.isna()].value_counts()","metadata":{"execution":{"iopub.status.busy":"2022-07-24T16:50:40.680847Z","iopub.execute_input":"2022-07-24T16:50:40.681969Z","iopub.status.idle":"2022-07-24T16:50:40.694076Z","shell.execute_reply.started":"2022-07-24T16:50:40.681905Z","shell.execute_reply":"2022-07-24T16:50:40.693116Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# That is how the tweets looked like when we tried appending the keywords\n(all_train_tweets.text + ' ' + all_train_tweets.keyword.fillna('')).loc[~all_train_tweets.keyword.isna()].to_list()[:5]","metadata":{"execution":{"iopub.status.busy":"2022-07-24T16:50:42.565927Z","iopub.execute_input":"2022-07-24T16:50:42.566631Z","iopub.status.idle":"2022-07-24T16:50:42.581033Z","shell.execute_reply.started":"2022-07-24T16:50:42.566598Z","shell.execute_reply":"2022-07-24T16:50:42.580017Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Location seems useful. Very messy, though.\nall_train_tweets.location.loc[~all_train_tweets.location.isna()].sort_values()","metadata":{"execution":{"iopub.status.busy":"2022-07-24T16:50:44.205859Z","iopub.execute_input":"2022-07-24T16:50:44.206942Z","iopub.status.idle":"2022-07-24T16:50:44.226820Z","shell.execute_reply.started":"2022-07-24T16:50:44.206873Z","shell.execute_reply":"2022-07-24T16:50:44.225991Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Preprocessing","metadata":{}},{"cell_type":"code","source":"def process_tweet(tweet, remove_stopwords=False, stem=False):\n    if tweet:\n        # Remove hyper-links\n        tweet = re.sub(r'https?:\\/\\/.*[\\r\\n]*', '', tweet)\n        # Remove hashtags\n        tweet = re.sub(r'#', '', tweet)\n        # Remove stock market tickers\n        tweet = re.sub(r'\\$\\w*', '', tweet)\n        # Remove old style tweet text RT\n        tweet = re.sub(r'^RT[\\s]+', '', tweet)\n        # Tokenize tweet\n        tokenizer = TweetTokenizer(preserve_case=False, strip_handles=True, reduce_len=True)\n        tweet = [word for word in tokenizer.tokenize(tweet) if word not in string.punctuation]\n        if remove_stopwords:\n            tweet = [word for word in tweet if word not in stopwords_english]\n        if stem:\n            stemmer = PorterStemmer()\n            tweet = [stemmer.stem(word) for word in tweet]\n    else:\n        tweet = ['__na__']\n\n    return tweet","metadata":{"execution":{"iopub.status.busy":"2022-07-24T16:50:48.568933Z","iopub.execute_input":"2022-07-24T16:50:48.570116Z","iopub.status.idle":"2022-07-24T16:50:48.580141Z","shell.execute_reply.started":"2022-07-24T16:50:48.570069Z","shell.execute_reply":"2022-07-24T16:50:48.578905Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def clean_dataset(dataset_df):\n    dataset_df['text_clean'] = dataset_df.text.apply(process_tweet, args=(True, False))\n    dataset_df['keyword_clean'] = dataset_df.keyword.fillna('').apply(process_tweet, args=(True, False))\n    dataset_df = pd.concat([dataset_df, tweet_geocode['location_clean']], axis=1, join='inner')\n    # dataset_df['location_clean'] = dataset_df.location.fillna('').apply(process_location)\n    return dataset_df\n    \nall_train_tweets = clean_dataset(all_train_tweets)\nall_test_tweets = clean_dataset(all_test_tweets)","metadata":{"execution":{"iopub.status.busy":"2022-07-24T16:50:50.343577Z","iopub.execute_input":"2022-07-24T16:50:50.343995Z","iopub.status.idle":"2022-07-24T16:50:52.843834Z","shell.execute_reply.started":"2022-07-24T16:50:50.343962Z","shell.execute_reply":"2022-07-24T16:50:52.842978Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"all_train_tweets","metadata":{"execution":{"iopub.status.busy":"2022-07-24T16:51:01.245600Z","iopub.execute_input":"2022-07-24T16:51:01.246015Z","iopub.status.idle":"2022-07-24T16:51:01.276887Z","shell.execute_reply.started":"2022-07-24T16:51:01.245983Z","shell.execute_reply":"2022-07-24T16:51:01.275608Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# The purpose of this is to show how long the sequences we are dealing with.\n# The longer the sequence, the tricker it is to capture the whole meaning of the tweet.\nall_train_tweets.text_clean.map(lambda t: len(t)).quantile([0.0, 0.25, .50, 0.75, 1.0])","metadata":{"execution":{"iopub.status.busy":"2022-07-24T16:51:05.479283Z","iopub.execute_input":"2022-07-24T16:51:05.480463Z","iopub.status.idle":"2022-07-24T16:51:05.498876Z","shell.execute_reply.started":"2022-07-24T16:51:05.480410Z","shell.execute_reply":"2022-07-24T16:51:05.497775Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Building vocabulary\ndef build_vocab(tweets):\n    vocab = {\n        '__pad__': 0,\n        '__unk__': 1,\n    }\n    for tweet in tweets:\n        for word in tweet:\n            if word not in vocab:\n                vocab[word] = len(vocab)\n    return vocab\n\n\n# Tweet to tensor\ndef tweet_to_tensor(tweet, vocab):\n    tweet = [vocab.get(token, vocab['__unk__']) for token in tweet]\n    return tweet\n\n\ndef prep_dataset(dataset_df, vocab):\n    dataset_df['text_clean'] = dataset_df.text_clean.apply(tweet_to_tensor, args=(vocab,))\n    dataset_df['keyword_clean'] = dataset_df.keyword_clean.apply(tweet_to_tensor, args=(vocab,))\n    # dataset_df['location_clean'] = dataset_df.location_clean.apply(tweet_to_tensor, args=(vocab,))\n    # dataset_df['input_clean'] = dataset_df.keyword_clean + dataset_df.location_clean\n\nvocab = build_vocab(all_train_tweets.text_clean.to_list())\nprep_dataset(all_train_tweets, vocab)\nprep_dataset(all_test_tweets, vocab)\n# prep_dataset(all_train_tweets, WORD_IDX)\n# prep_dataset(all_test_tweets, WORD_IDX)","metadata":{"execution":{"iopub.status.busy":"2022-07-24T16:51:09.107332Z","iopub.execute_input":"2022-07-24T16:51:09.108065Z","iopub.status.idle":"2022-07-24T16:51:09.191730Z","shell.execute_reply.started":"2022-07-24T16:51:09.108031Z","shell.execute_reply":"2022-07-24T16:51:09.190861Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Train/Validation split\nall_pos_train = all_train_tweets.loc[all_train_tweets.target == 1]\nall_neg_train = all_train_tweets.loc[all_train_tweets.target == 0]\n\npos_cut_idx = int(all_pos_train.shape[0] * (1 - VAL_PCT))\npos_val = all_pos_train.iloc[pos_cut_idx:]\npos_train = all_pos_train.iloc[:pos_cut_idx]\n\nneg_cut_idx = int(all_neg_train.shape[0] * (1 - VAL_PCT))\nneg_val = all_neg_train.iloc[neg_cut_idx:]\nneg_train = all_neg_train.iloc[:neg_cut_idx] \n\nall_train = pd.concat([pos_train, neg_train])\nall_val = pd.concat([pos_val, neg_val])","metadata":{"execution":{"iopub.status.busy":"2022-07-24T16:51:11.585476Z","iopub.execute_input":"2022-07-24T16:51:11.585852Z","iopub.status.idle":"2022-07-24T16:51:11.603798Z","shell.execute_reply.started":"2022-07-24T16:51:11.585815Z","shell.execute_reply":"2022-07-24T16:51:11.602941Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# We tried different ways to train the models. Not all parameters are currently used\ndef data_generator(text_pos, text_neg, keyword_pos, keyword_neg, loc_pos, loc_neg, batch_size, vocab, loop=False):\n    len_pos = len(text_pos)\n    len_neg = len(text_neg)\n    \n    pos_idx_lines =  list(range(len_pos))\n    neg_idx_lines = list(range(len_neg))\n    \n    pos_idx = 0\n    neg_idx = 0\n    \n    n_to_take = batch_size // 2\n    \n    random.shuffle(pos_idx_lines)\n    random.shuffle(neg_idx_lines)\n    \n    stop = False\n    \n    while not stop:\n        batch_text = []\n        batch_keyword = []\n        batch_loc = []\n        targets = []\n        max_len_text = 0\n        max_len_keyword = 0\n        max_len_loc = 0\n        for i in range(n_to_take):\n            if pos_idx >= len_pos or neg_idx >= len_neg:\n                if not loop:\n                    stop = True\n                    break\n                if pos_idx >= len_pos:\n                    pos_idx = 0\n                    random.shuffle(pos_idx_lines)\n                if neg_idx >= len_neg:\n                    neg_idx = 0\n                    random.shuffle(neg_idx_lines)\n                    \n            # Tweet body\n            pos_text = text_pos[pos_idx]\n            batch_text.append(pos_text)\n            if len(pos_text) > max_len_text:\n                max_len_text = len(pos_text)\n            targets.append(1)\n                \n            neg_text = text_neg[neg_idx]\n            batch_text.append(neg_text)\n            if len(neg_text) > max_len_text:\n                max_len_text = len(neg_text)\n            targets.append(0)\n            \n            # Keyword\n            if keyword_pos:\n                pos_keyword = keyword_pos[pos_idx]\n                batch_keyword.append(pos_keyword)\n                if len(pos_keyword) > max_len_keyword:\n                    max_len_keyword = len(pos_keyword)\n                    \n                neg_keyword = keyword_neg[neg_idx]\n                batch_keyword.append(neg_keyword)\n                if len(neg_keyword) > max_len_keyword:\n                    max_len_keyword = len(neg_keyword)\n            \n            # Location\n            if loc_pos:\n                pos_loc = loc_pos[pos_idx]\n                batch_loc.append(pos_loc)\n                # if len(pos_loc) > max_len_loc:\n                #     max_len_loc = len(pos_loc)\n\n                neg_loc = loc_neg[neg_idx]\n                batch_loc.append(neg_loc)\n                # if len(neg_loc) > max_len_loc:\n                #     max_len_loc = len(neg_loc)\n\n            pos_idx += 1\n            neg_idx += 1\n                \n        if stop:\n            break\n            \n        pos_idx += n_to_take\n        neg_idx += n_to_take\n        \n        # padding\n        for elem in batch_text:\n            elem += [vocab['__pad__']] * (max_len_text - len(elem))\n        for elem in batch_keyword:\n            elem += [vocab['__pad__']] * (max_len_keyword - len(elem))\n        # for elem in batch_loc:\n        #     elem += [vocab['__pad__']] * (max_len_loc - len(elem))\n            \n        # We do not use them here, but they are expected by Trax\n        example_weights = np.array([1] * (n_to_take * 2))\n            \n        ret_vals = (np.array(batch_text),)\n        if keyword_pos:\n            ret_vals += (np.array(batch_keyword),)\n        if loc_pos:\n            ret_vals += (np.array(batch_loc),)\n        yield ret_vals + (np.array(targets), example_weights)","metadata":{"execution":{"iopub.status.busy":"2022-07-24T16:51:13.617833Z","iopub.execute_input":"2022-07-24T16:51:13.618866Z","iopub.status.idle":"2022-07-24T16:51:13.638719Z","shell.execute_reply.started":"2022-07-24T16:51:13.618814Z","shell.execute_reply":"2022-07-24T16:51:13.637380Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_generator(batch_size, train_text_pos, train_text_neg, train_keyword_pos, train_keyword_neg, train_loc_pos, train_loc_neg, vocab):\n    return data_generator(\n        text_pos=train_text_pos,\n        text_neg=train_text_neg,\n        keyword_pos=train_keyword_pos,\n        keyword_neg=train_keyword_neg,\n        loc_pos=train_loc_pos,\n        loc_neg=train_loc_neg,\n        batch_size=batch_size,\n        vocab=vocab,\n        loop=True\n    )\ndef val_generator(batch_size, val_text_pos, val_text_neg, val_keyword_pos, val_keyword_neg, val_loc_pos, val_loc_neg, vocab):\n    return data_generator(\n        text_pos=val_text_pos,\n        text_neg=val_text_neg,\n        keyword_pos=val_keyword_pos,\n        keyword_neg=val_keyword_neg,\n        loc_pos=val_loc_pos,\n        loc_neg=val_loc_neg,\n        batch_size=batch_size,\n        vocab=vocab,\n        loop=True\n    )","metadata":{"execution":{"iopub.status.busy":"2022-07-24T16:51:16.895050Z","iopub.execute_input":"2022-07-24T16:51:16.895805Z","iopub.status.idle":"2022-07-24T16:51:16.903389Z","shell.execute_reply.started":"2022-07-24T16:51:16.895759Z","shell.execute_reply":"2022-07-24T16:51:16.902560Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_train_eval_tasks(\n    train_text_pos, train_text_neg, train_keyword_pos, train_keyword_neg, train_loc_pos, train_loc_neg,\n    val_text_pos, val_text_neg, val_keyword_pos, val_keyword_neg, val_loc_pos, val_loc_neg,\n    batch_size, vocab\n):\n    train_task = training.TrainTask(\n        labeled_data=train_generator(\n            batch_size=batch_size,\n            train_text_pos=train_text_pos,\n            train_text_neg=train_text_neg,\n            train_keyword_pos=train_keyword_pos,\n            train_keyword_neg=train_keyword_neg,\n            train_loc_pos=train_loc_pos,\n            train_loc_neg=train_loc_neg,\n            vocab=vocab,\n        ),\n        loss_layer=tl.CrossEntropyLoss(),\n        optimizer=optimizers.Adam(0.001),\n        n_steps_per_checkpoint=10,\n    )\n    eval_task = training.EvalTask(\n        labeled_data=val_generator(\n            batch_size=batch_size,\n            val_text_pos=val_text_pos,\n            val_text_neg=val_text_neg,\n            val_keyword_pos=val_keyword_pos,\n            val_keyword_neg=val_keyword_neg,\n            val_loc_pos=val_loc_pos,\n            val_loc_neg=val_loc_neg,\n            vocab=vocab,\n        ),\n        metrics=[\n            tl.CrossEntropyLoss(),\n            tl.Accuracy(),\n            tl.MacroAveragedFScore(),\n        ]\n    )\n    \n    return train_task, eval_task","metadata":{"execution":{"iopub.status.busy":"2022-07-24T16:51:17.594103Z","iopub.execute_input":"2022-07-24T16:51:17.594673Z","iopub.status.idle":"2022-07-24T16:51:17.603922Z","shell.execute_reply.started":"2022-07-24T16:51:17.594641Z","shell.execute_reply":"2022-07-24T16:51:17.602868Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BATCH_SIZE = 32\ntrain_task, eval_task = get_train_eval_tasks(\n    train_text_pos=pos_train.text_clean.to_list(),\n    train_text_neg=neg_train.text_clean.to_list(),\n    train_keyword_pos=pos_train.keyword_clean.to_list(),\n    train_keyword_neg=neg_train.keyword_clean.to_list(),\n    train_loc_pos=pos_train.location_clean.to_list(),\n    train_loc_neg=neg_train.location_clean.to_list(),\n    val_text_pos=pos_val.text_clean.to_list(),\n    val_text_neg=neg_val.text_clean.to_list(),\n    val_keyword_pos=pos_val.keyword_clean.to_list(),\n    val_keyword_neg=neg_val.keyword_clean.to_list(),\n    val_loc_pos=pos_val.location_clean.to_list(),\n    val_loc_neg=neg_val.location_clean.to_list(),\n    batch_size=BATCH_SIZE,\n    vocab=vocab,\n)","metadata":{"execution":{"iopub.status.busy":"2022-07-24T16:51:18.446002Z","iopub.execute_input":"2022-07-24T16:51:18.446389Z","iopub.status.idle":"2022-07-24T16:51:18.916301Z","shell.execute_reply.started":"2022-07-24T16:51:18.446359Z","shell.execute_reply":"2022-07-24T16:51:18.915154Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Building the model","metadata":{}},{"cell_type":"code","source":"# This is where we treid to create a non-trainable Trax Embedding layer.\n# Did not really work\nclass NonTrainableEmbedding(tl.Embedding):\n    @property\n    def has_backward(self):\n        return True\n\n    def backward(self, inputs, output, grad, weights, state, new_state, rng):\n        return np.zeros(inputs.shape), np.zeros(weights.shape)","metadata":{"execution":{"iopub.status.busy":"2022-07-24T16:51:21.272960Z","iopub.execute_input":"2022-07-24T16:51:21.274032Z","iopub.status.idle":"2022-07-24T16:51:21.279692Z","shell.execute_reply.started":"2022-07-24T16:51:21.273986Z","shell.execute_reply":"2022-07-24T16:51:21.278839Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# I also tried adding an LSTM layer and selecting the last word of a seqeuce.\n# Because the sequence are padded this function is not correct. I should have\n# selected the last non-pad word in each sequence.\ndef select_last(seq):\n    return seq[:,-1,:]\n    \n# Because we are training three different branches, we need to combine their\n# scores somehow\ndef calc_score(text_score, keyword_score, location_score):\n    return text_score * 0.4 + keyword_score * 0.4 + location_score * 0.2\n\n# In one experiment I tried to use the embedding matrix as is without encapsulation\n# inside an Embedding layer\ndef get_word_embed(sequences):\n    for seq in sequences:\n        for i, word_idx in enumerate(seq):\n            embed = EMBED_MATRIX.get(IDX_WORD[word_idx.astype('int')], [0] * EMBED_SIZE)\n            print(f'>>> {word_idx} -> {embed}')\n            seq[i] = np.array()\n    return sequences\n\n# Sometimes shapes get really confusing\ndef print_shape(seq):\n    print(f'>>> seq {seq.shape}')\n    return seq\n\n\n    \ndef classifier(vocab_size, embedding_dim):\n# def classifier(embed_file, vocab_size, embed_size, batch_size):\n    embed_l1 = tl.Embedding(vocab_size, embedding_dim)\n    # embed_l1 = NonTrainableEmbedding(vocab_size, embed_size)\n    # embed_l1.init_from_file(\n    #     file_name=embed_file,\n    #     weights_only=True,\n    #     input_signature=trax.shapes.signature(np.empty(shape=(batch_size, 26))),\n    # )\n    # embed_l2 = tl.Embedding(vocab_size, embed_size)\n    # embed_l2.init_from_file(\n    #    file_name=embed_file,\n    #    weights_only=True,\n    #    input_signature=trax.shapes.signature(np.empty(shape=(batch_size, 26))),\n    # )\n    \n    # dense_l1 = tl.Dense(n_units=2)\n    # dense_l1.init_weights_and_state(trax.shapes.signature(np.empty(shape=(batch_size, embed_size))))\n    # dense_l2 = tl.Dense(n_units=2)\n    # dense_l2.init_weights_and_state(trax.shapes.signature(np.empty(shape=(batch_size, embed_size))))\n    \n    return tl.Serial(\n        tl.Parallel(\n            tl.Serial(\n                # tl.Embedding(\n                #    vocab_size=vocab_size,\n                #    d_feature=embedding_dim_text,\n                # ),\n                embed_l1,\n                tl.Mean(axis=1),\n                tl.Dense(n_units=100),\n                tl.Dense(n_units=2),\n                # dense_l1,\n                tl.LogSoftmax(),\n            ),\n            tl.Serial(\n                # tl.Embedding(\n                #    vocab_size=vocab_size,\n                #    d_feature=embedding_dim_kw,\n                # ),\n                embed_l1,\n                tl.Mean(axis=1),\n                tl.Dense(n_units=100),\n                tl.Dense(n_units=2),\n                # dense_l2,\n                tl.LogSoftmax(),\n            ),\n            tl.Serial(\n                tl.Dense(n_units=5),\n                tl.Dense(n_units=2),\n                # dense_l2,\n                tl.LogSoftmax(),\n            ),\n        ),\n        tl.Fn('calc_score', calc_score),\n    )","metadata":{"execution":{"iopub.status.busy":"2022-07-24T16:52:21.676798Z","iopub.execute_input":"2022-07-24T16:52:21.677686Z","iopub.status.idle":"2022-07-24T16:52:21.692140Z","shell.execute_reply.started":"2022-07-24T16:52:21.677643Z","shell.execute_reply":"2022-07-24T16:52:21.691130Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Embedding dimension. Something to experiemnt with.\nEMBED_DIM_TEXT = 265\n# EMBED_DIM_KW = 128\n\nmodel = classifier(len(vocab), EMBED_DIM_TEXT)\n# model = classifier(PROC_EMBED_FILE, len(WORD_IDX), EMBED_SIZE, BATCH_SIZE)\nmodel","metadata":{"execution":{"iopub.status.busy":"2022-07-24T16:52:24.476087Z","iopub.execute_input":"2022-07-24T16:52:24.476500Z","iopub.status.idle":"2022-07-24T16:52:24.485700Z","shell.execute_reply.started":"2022-07-24T16:52:24.476466Z","shell.execute_reply":"2022-07-24T16:52:24.484608Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training the model","metadata":{}},{"cell_type":"code","source":"def train_model(model, train_task, eval_task, n_steps, output_dir):\n    training_loop = training.Loop(\n        model=model,\n        tasks=train_task,\n        eval_tasks=eval_task,\n        output_dir=output_dir,\n    )\n    \n    training_loop.run(n_steps=n_steps)\n    \n    return training_loop","metadata":{"execution":{"iopub.status.busy":"2022-07-24T16:52:33.698374Z","iopub.execute_input":"2022-07-24T16:52:33.698773Z","iopub.status.idle":"2022-07-24T16:52:33.704778Z","shell.execute_reply.started":"2022-07-24T16:52:33.698744Z","shell.execute_reply":"2022-07-24T16:52:33.703724Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Empty the model directory, otherwise old checkpoints might fail to load and training would fail\ndef delete_model():\n    for filename in os.listdir(MODEL_DIR):\n        file_path = os.path.join(MODEL_DIR, filename)\n        try:\n            if os.path.isfile(file_path) or os.path.islink(file_path):\n                os.unlink(file_path)\n            elif os.path.isdir(file_path):\n                shutil.rmtree(file_path)\n        except Exception as e:\n            print('Failed to delete %s. Reason: %s' % (file_path, e))","metadata":{"execution":{"iopub.status.busy":"2022-07-24T16:53:09.650566Z","iopub.execute_input":"2022-07-24T16:53:09.651204Z","iopub.status.idle":"2022-07-24T16:53:09.659176Z","shell.execute_reply.started":"2022-07-24T16:53:09.651164Z","shell.execute_reply":"2022-07-24T16:53:09.657863Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"training_loop = train_model(\n    model=model,\n    train_task=train_task,\n    eval_task=[eval_task],\n    n_steps=150,\n    output_dir=MODEL_DIR,\n)","metadata":{"execution":{"iopub.status.busy":"2022-07-24T16:53:21.883494Z","iopub.execute_input":"2022-07-24T16:53:21.883886Z","iopub.status.idle":"2022-07-24T16:55:02.450767Z","shell.execute_reply.started":"2022-07-24T16:53:21.883856Z","shell.execute_reply":"2022-07-24T16:55:02.449668Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"hist_acc = training_loop.history.to_dict()['_values']['eval']['metrics/Accuracy']\nhist_f1 = training_loop.history.to_dict()['_values']['eval']['metrics/MacroAveragedFScore']\nsteps = [t[0] for t in hist_f1]\nvals_f1 = [t[1] for t in hist_f1]\nvals_acc = [t[1] for t in hist_acc]\nimport statistics as st\navg_f1 = st.mean(vals_f1)\navg_acc = st.mean(vals_acc)\n\nfig, ax = plt.subplots(figsize=(10, 8))\nax.plot(steps, vals_f1);\nax.plot(steps, [avg_f1] * len(vals_f1), linestyle=\"-.\")\nax.plot(steps, [avg_acc] * len(vals_acc), linestyle=\":\")\nprint(f'>>> Average F1 {avg_f1}')\nprint(f'>>> Average Accuracy {avg_acc}')","metadata":{"execution":{"iopub.status.busy":"2022-07-24T16:59:22.692892Z","iopub.execute_input":"2022-07-24T16:59:22.693432Z","iopub.status.idle":"2022-07-24T16:59:22.968393Z","shell.execute_reply.started":"2022-07-24T16:59:22.693396Z","shell.execute_reply":"2022-07-24T16:59:22.967340Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Testing the model","metadata":{}},{"cell_type":"code","source":"# Prediction\nmax_len_text = all_test_tweets.text_clean.map(len).max().item()\nmax_len_kw = all_test_tweets.keyword_clean.map(len).max().item()\n\ndef pad(tweet, max_len):\n    return tweet + ([0] * (max_len - len(tweet)))\n\nall_test_tweets['text_clean'] = all_test_tweets.text_clean.apply(pad, args=(max_len_text,))\nall_test_tweets['keyword_clean'] = all_test_tweets.keyword_clean.apply(pad, args=(max_len_kw,))\n\npreds = training_loop.eval_model((np.array(all_test_tweets.text_clean.to_list()), np.array(all_test_tweets.keyword_clean.to_list()), np.array(all_test_tweets.location_clean.to_list())))\ntarget = np.array([pred[1] > pred[0] for pred in preds]).astype(np.integer)\nall_test_tweets['target'] = target","metadata":{"execution":{"iopub.status.busy":"2022-07-24T16:59:31.105496Z","iopub.execute_input":"2022-07-24T16:59:31.105901Z","iopub.status.idle":"2022-07-24T16:59:41.509712Z","shell.execute_reply.started":"2022-07-24T16:59:31.105870Z","shell.execute_reply":"2022-07-24T16:59:41.508564Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Output\nall_test_tweets[['target']].to_csv(f'{OUTPUT_DIR}/submission.csv')","metadata":{"execution":{"iopub.status.busy":"2022-07-24T17:02:02.177585Z","iopub.execute_input":"2022-07-24T17:02:02.178079Z","iopub.status.idle":"2022-07-24T17:02:02.192358Z","shell.execute_reply.started":"2022-07-24T17:02:02.178044Z","shell.execute_reply":"2022-07-24T17:02:02.191558Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}