{"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":"<h1 style=\"text-align: center; font-family: Verdana; font-size: 32px; font-style: normal; font-weight: bold; text-decoration: none; text-transform: none; font-variant: small-caps; letter-spacing: 3px; color: #FF1493; background-color: #ffffff;\">Bristol-Myers Squibb – Molecular Translation</h1>\n<h2 style=\"text-align: center; font-family: Verdana; font-size: 24px; font-style: normal; font-weight: bold; text-decoration: underline; text-transform: none; letter-spacing: 2px; color: navy; background-color: #ffffff;\">TFRecord Creation</h2>\n<h5 style=\"text-align: center; font-family: Verdana; font-size: 12px; font-style: normal; font-weight: bold; text-decoration: None; text-transform: none; letter-spacing: 1px; color: black; background-color: #ffffff;\">CREATED BY: DARIEN SCHETTLER</h5>\n\n---\n\n<br><br>\n\n<font color=\"purple\" style=\"font-weight: bold;\">CHANGE LOG</font>\n\n---\n\n* **v1 - `Working Draft`**\n* **v2 - `384x384x1 - No Rotation - Raw Labels - Preprocessing (inversion, pad2square)`**\n* **v3-9 - `Working Draft`**\n* **v10 - `384x384x1 - Yes Rotation - Tokenized - Preprocessing (inversion, pad2square) - Just Train/Val `**\n* **v11 - `128x128x1 - Yes Rotation - Tokenized - Preprocessing (inversion, pad2square) - Train/Val/Test `**\n* **v12 - `256x256x1 - Yes Rotation - Tokenized - Preprocessing (inversion, pad2square) - Train/Val/Test [CURRENT]`**\n* **v13-15 - `Working Draft`**\n\n---\n\n<br>","metadata":{}},{"cell_type":"markdown","source":"<p id=\"toc\"></p>\n\n<br><br>\n\n<h1 style=\"font-family: Verdana; font-size: 24px; font-style: normal; font-weight: bold; text-decoration: none; text-transform: none; letter-spacing: 3px; color: navy; background-color: #ffffff;\">TABLE OF CONTENTS</h1>\n\n---\n\n<h3 style=\"text-indent: 10vw; font-family: Verdana; font-size: 20px; font-style: normal; font-weight: normal; text-decoration: none; text-transform: none; letter-spacing: 2px; color: navy; background-color: #ffffff;\"><a href=\"#imports\">0&nbsp;&nbsp;&nbsp;&nbsp;IMPORTS</a></h3>\n\n---\n\n<h3 style=\"text-indent: 10vw; font-family: Verdana; font-size: 20px; font-style: normal; font-weight: normal; text-decoration: none; text-transform: none; letter-spacing: 2px; color: navy; background-color: #ffffff;\"><a href=\"#background_information\">1&nbsp;&nbsp;&nbsp;&nbsp;BACKGROUND INFORMATION</a></h3>\n\n---\n\n<h3 style=\"text-indent: 10vw; font-family: Verdana; font-size: 20px; font-style: normal; font-weight: normal; text-decoration: none; text-transform: none; letter-spacing: 2px; color: navy; background-color: #ffffff;\"><a href=\"#setup\">2&nbsp;&nbsp;&nbsp;&nbsp;SETUP</a></h3>\n\n---\n\n<h3 style=\"text-indent: 10vw; font-family: Verdana; font-size: 20px; font-style: normal; font-weight: normal; text-decoration: none; text-transform: none; letter-spacing: 2px; color: navy; background-color: #ffffff;\"><a href=\"#helper_functions\">3&nbsp;&nbsp;&nbsp;&nbsp;HELPER FUNCTIONS</a></h3>\n\n---\n\n<h3 style=\"text-indent: 10vw; font-family: Verdana; font-size: 20px; font-style: normal; font-weight: normal; text-decoration: none; text-transform: none; letter-spacing: 2px; color: navy; background-color: #ffffff;\"><a href=\"#dataset_preparation\">4&nbsp;&nbsp;&nbsp;&nbsp;PREPARE THE DATASET</a></h3>\n\n---","metadata":{}},{"cell_type":"markdown","source":"<h1 style=\"font-family: Verdana; font-size: 24px; font-style: normal; font-weight: bold; text-decoration: none; text-transform: none; letter-spacing: 3px; background-color: #ffffff; color: navy;\" id=\"imports\">0&nbsp;&nbsp;IMPORTS&nbsp;&nbsp;&nbsp;&nbsp;<a href=\"#toc\">&#10514;</a></h1>","metadata":{}},{"cell_type":"code","source":"print(\"\\n... IMPORTS STARTING ...\\n\")\n# print(\"\\n\\tVERSION INFORMATION\")\n\n# # Parallel application across pandas apply operation\n# !pip install pandarallel pickle5\n# from pandarallel import pandarallel\n# pandarallel.initialize()\n\n# Machine Learning and Data Science Imports\nimport tensorflow as tf; print(f\"\\t\\t– TENSORFLOW VERSION: {tf.__version__}\");\nimport tensorflow_addons as tfa; print(f\"\\t\\t– TENSORFLOW ADDONS VERSION: {tfa.__version__}\");\n# import pandas as pd; pd.options.mode.chained_assignment = None;\nimport numpy as np; print(f\"\\t\\t– NUMPY VERSION: {np.__version__}\");\n\n# # Built In Imports\n# from kaggle_datasets import KaggleDatasets\n# from collections import Counter\n# from datetime import datetime\n# import multiprocessing as mp\n# from glob import glob\n# import warnings\n# import requests\n# import imageio\n# import IPython\n# import urllib\n# import zipfile\n# import pickle\n# import random\n# import shutil\n# import string\n# import math\n# import time\n# import gzip\n# import ast\n# import io\nimport os\n# import gc\n# import re\n\n# # Visualization Imports\n# from matplotlib.colors import ListedColormap\n# import matplotlib.patches as patches\n# import plotly.graph_objects as go\n# import matplotlib.pyplot as plt\nfrom tqdm.notebook import tqdm; tqdm.pandas();\n# import plotly.express as px\n# import seaborn as sns\n# from PIL import Image\n# import matplotlib; print(f\"\\t\\t– MATPLOTLIB VERSION: {matplotlib.__version__}\");\n# import plotly\n# import PIL\n# import cv2\nimport math\n\n\nprint(\"\\n\\n... IMPORTS COMPLETE ...\\n\")\n\ngpus = tf.config.experimental.list_physical_devices('GPU')\nif gpus:\n    try:\n        # Currently, memory growth needs to be the same across GPUs\n        for gpu in gpus:\n            tf.config.experimental.set_memory_growth(gpu, True)\n        logical_gpus = tf.config.experimental.list_logical_devices('GPU')\n        print(len(gpus), \"... Physical GPUs,\", len(logical_gpus), \"Logical GPUs ...\\n\")\n    except RuntimeError as e:\n        # Memory growth must be set before GPUs have been initialized\n        print(e)\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h1 style=\"font-family: Verdana; font-size: 24px; font-style: normal; font-weight: bold; text-decoration: none; text-transform: none; letter-spacing: 3px; color: navy; background-color: #ffffff;\" id=\"background_information\">1&nbsp;&nbsp;BACKGROUND INFORMATION&nbsp;&nbsp;&nbsp;&nbsp;<a href=\"#toc\">&#10514;</a></h1>\n\n<br><b style=\"text-decoration: underline; font-family: Verdana; text-transform: uppercase;\">POSSIBLE FLAG EXPLORATION</b>\n\nTBD","metadata":{}},{"cell_type":"markdown","source":"<h1 style=\"font-family: Verdana; font-size: 24px; font-style: normal; font-weight: bold; text-decoration: none; text-transform: none; letter-spacing: 3px; color: navy; background-color: #ffffff;\" id=\"setup\">2&nbsp;&nbsp;SETUP&nbsp;&nbsp;&nbsp;&nbsp;<a href=\"#toc\">&#10514;</a></h1>","metadata":{}},{"cell_type":"markdown","source":"<h3 style=\"font-family: Verdana; font-size: 20px; font-style: normal; font-weight: normal; text-decoration: none; text-transform: none; letter-spacing: 2px; color: navy; background-color: #ffffff;\">2.1  BASIC SETUP</h3>","metadata":{}},{"cell_type":"code","source":"# # Only required if `APPLY_TOKENIZATION=True`\n# # TOKEN_LIST = ['<PAD>', '<START>', '<END>']\n# # TOKEN_LIST += ['\\[SiH2+\\]', '\\[SiH2-\\]', '\\[SiH2\\]', '\\[SiH3\\]', '\\[SiH\\]', '\\[Si+\\]', '\\[Si-\\]', '\\[Si@@\\]', '\\[Si@\\]', '\\[Si\\]']\n# # TOKEN_LIST += ['\\[NH+\\]', '\\[NH\\]', '\\[N+\\]', '\\[N-\\]', '\\[N\\]']\n# # TOKEN_LIST += ['\\[P+\\]', '\\[P\\]', '\\[S\\]', '\\%10', '\\[B-\\]','\\[B+\\]', '\\[B\\]', '\\[Cl+3\\]', '\\[PH\\]', '\\[I\\]']\n# # TOKEN_LIST += ['\\[P@@\\]', '\\[P@\\]', '\\[CH2+\\]', '\\[CH2\\]', '\\[CH+\\]', '\\[CH-\\]', '\\[CH\\]', '\\[C+\\]', '\\[C-\\]', '\\[C\\]', '\\[O-\\]', '\\[O+\\]', '\\[O\\]']\n# # TOKEN_LIST += ['\\[S@@\\]', '\\[S@\\]', '\\[2H\\]', '\\[C@@H\\]', '\\[C@@\\]', '\\[C@H\\]', '\\[C@\\]']\n# # TOKEN_LIST += ['Br', 'Cl', 'B', 'P', 'I', 'O', 'N', 'S', 'F', 'C', '9', '8', '7', '6', '5', '4', '3', '2', '1',  '\\\\\\\\', '#', '\\/', '=', ':', '\\(', '\\)']\n# # # START_TOKEN = tf.constant(TOKEN_LIST.index(), dtype=tf.uint8)\n# # END_TOKEN = tf.constant(TOKEN_LIST.index(\"<END>\"), dtype=tf.uint8)\n# # PAD_TOKEN = tf.constant(TOKEN_LIST.index(\"<PAD>\"), dtype=tf.uint8)\n\n# # Define the root and data directories\n# ROOT_DIR = \"/kaggle/input\"\n# DATA_DIR = os.path.join(ROOT_DIR, \"bms-molecular-translation\")\n# TRAIN_DIR = os.path.join(DATA_DIR, \"train\")\n# TEST_DIR = os.path.join(DATA_DIR, \"test\")\n\n# TRAIN_CSV_PATH = \"/kaggle/input/train-labels-v5/train_labels_v5.pickle\"\n# SS_CSV_PATH = \"/kaggle/input/bms-csvs-w-extra-metadata/sample_submission_w_extra.csv\"\n\n# import pickle5 as pickle\n# with open(TRAIN_CSV_PATH, 'rb') as fh:\n# train_df = pd.read_pickle(TRAIN_CSV_PATH)\n    \n# train_df.dropna(inplace=True)\n# train_df.roi_bbox = train_df.roi_bbox.progress_apply(lambda x: ast.literal_eval(x))\n\n# print(\"\\n... TRAIN DATAFRAME W/ PATHS ...\\n\")\n# display(train_df)\n\n# ss_df = pd.read_csv(SS_CSV_PATH)\n# ss_df.roi_bbox = ss_df.roi_bbox.progress_apply(lambda x: ast.literal_eval(x))\n\n# print(\"\\n... SUBMISSION DATAFRAME ...\\n\")\n# display(ss_df)\n\n# def get_new_ar(row):\n#     if row.crop_width<row.crop_height:\n#         return row.crop_height/row.crop_width\n#     else:\n#         return row.aspect_ratio\n\n# ss_df.progress_apply(lambda x: get_new_ar(x), axis=1).describe()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 style=\"font-family: Verdana; font-size: 20px; font-style: normal; font-weight: normal; text-decoration: none; text-transform: none; letter-spacing: 2px; color: navy; background-color: #ffffff;\">2.2 USER INPUT VARIABLES</h3>","metadata":{}},{"cell_type":"code","source":"ROOT_DIR = \"/kaggle/input\"\nTRAIN_DIR = \"gs://from_aws_0001/subset_1_train_1a\"\nTEST_DIR = os.path.join(ROOT_DIR, \"02-selfies-tfrecord-creation\")\n\nselfies_train_dir = os.path.join(ROOT_DIR, \"decimer-subset-1/subset_1_train_1_SELFIES.txt\")\nselfies_test_dir =  os.path.join(ROOT_DIR, \"decimer-subset-1/subset_1_test_2_SELFIES.txt\")\n\npattern_train_dir =  os.path.join(ROOT_DIR, \"decimer-subset-1/subset_1_train_1_morgan.fps\")\npattern_test_dir =  os.path.join(ROOT_DIR, \"decimer-subset-1/subset_1_test_2_morgan.fps\")\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 style=\"font-family: Verdana; font-size: 20px; font-style: normal; font-weight: normal; text-decoration: none; text-transform: none; letter-spacing: 2px; color: navy; background-color: #ffffff;\">2.3 AUTO DETECTED VARIABLES</h3>","metadata":{}},{"cell_type":"code","source":"# Smile RELATED\n# TOK2INT = {c.strip(\"\\\\\"):i for i,c in enumerate(TOKEN_LIST)}\n# TOK2INT = {}\n# for i, c in enumerate(TOKEN_LIST):\n#     if not c == \"\\\\\\\\\":\n#         TOK2INT[c.replace(\"\\\\\", \"\")] = i\n#     else:\n#         TOK2INT[\"\\\\\"] = i\n\n# INT2TOK = {v:k for k,v in TOK2INT.items()}\n# # +2 for starting and ending tokens\n# MAX_LEN = train_df.smile.progress_apply(lambda x: len(re.findall(\"|\".join(TOKEN_LIST), x))).max() + 2\n# #     - Using half of the actual max length to accelerate training\n# FILTER_ON_MAX_LEN=True\n# # MAX_LEN = int((train_df.smile_token_len.max()+1)//2) # +1 is for the end token\n# VOCAB_LEN = len(INT2TOK)\n\n# # IMG RELATED\n# N_CHANNELS = IMG_SHAPE[-1]\n# IMG_SIZE = IMG_SHAPE[:2]\n\n# # Split dataframe\n# val_df = train_df[:N_VAL].reset_index(drop=True)\n# train_df = train_df[N_VAL:].reset_index(drop=True)\n\n# # For TFRecord sharding - very rough estimations\n# N_EX_PER_REC = (20000 if np.product(IMG_SHAPE)>400000 else 40000)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h1 style=\"font-family: Verdana; font-size: 24px; font-style: normal; font-weight: bold; text-decoration: none; text-transform: none; letter-spacing: 3px; color: navy; background-color: #ffffff;\" id=\"helper_functions\">3&nbsp;&nbsp;HELPER FUNCTIONS&nbsp;&nbsp;&nbsp;&nbsp;<a href=\"#toc\">&#10514;</a></h1>","metadata":{}},{"cell_type":"code","source":"\ndef tf_load_image(path, img_size=(224,224), tile_to_3_channel=True, invert=False, rotate_trick=False, repair_images=False):\n    \"\"\" Load an image with the correct size and shape \"\"\"\n    img = decode_img(tf.io.read_file(path), img_size, n_channels=1, invert=invert, rotate_trick=rotate_trick, repair_images=repair_images)\n    \n    if tile_to_3_channel:\n        return tf.tile(img, tf.constant((1, 1, 3), dtype=tf.int32))\n    else:\n        return img\n    \n    \ndef rotate_trick_fn(a):\n    \"\"\" Pad a tensor array `a` evenly until it is a square \"\"\"\n    h_src = tf.shape(a)[0]\n    w_src = tf.shape(a)[1] \n            \n    if h_src>w_src: # pad width\n        a = tf.image.rot90(a)\n\n    return a\n    \n    \ndef decode_img(img, img_size=(224,224),resize=False, n_channels=1, invert=False, rotate_trick=False, repair_images=False):\n    \n    \"\"\" Decode the image by utilizing TF ... pad to square ... and resize \"\"\"\n    \n    # convert the compressed string to a 3D uint8 tensor\n    img = tf.image.decode_png(img, channels=n_channels)\n    \n    if invert:\n        img = tf.ones_like(img, dtype=tf.uint8)*255-img\n        constant_pad=0\n    else:\n        constant_pad=255\n        \n    if repair_images:\n        pass\n        \n    # rotate trick\n    if rotate_trick:\n        img = rotate_trick_fn(img)\n    if resize:\n        img = tf.image.resize(img, img_size, method=\"bilinear\")\n    return tf.cast(img, tf.uint8)\n\n\ndef fpstolist(fp):\n    fp_int = [int(f,16) for f in fp]\n    fp_b = ['{:04b}'.format(f) for f in fp_int]\n    fp_bin=''.join((x for x in fp_b))\n    # List\n    fp_list = [int(t) for t in fp_bin]\n    return  fp_list\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h1 style=\"font-family: Verdana; font-size: 24px; font-style: normal; font-weight: bold; text-decoration: none; text-transform: none; letter-spacing: 3px; color: navy; background-color: #ffffff;\" id=\"dataset_preparation\">4&nbsp;&nbsp;PREPARE THE DATASET&nbsp;&nbsp;&nbsp;&nbsp;<a href=\"#toc\">&#10514;</a></h1>\n\nIn this section we prepare a subset of the dataset for modelling","metadata":{}},{"cell_type":"code","source":"\ntest_selfies = []\ntest_image_id =[]\nwith open(selfies_test_dir) as file:\n    lines = file.readlines()\n    for line in tqdm(lines):\n        tokens = line.split(\",\")\n        image_id = str(tokens[0]) \n\n        try:\n            Selfies_1 = tokens[1].split('[')\n            Selfies=['['+x.strip() for x in Selfies_1 if x.strip() != '']\n            Selfies= \" \".join(str(i) for i in Selfies)\n            SMILES_ = \"<start> \" + Selfies + \" <end>\"\n            test_selfies.append(SMILES_)\n            test_image_id.append(image_id)\n        except IndexError as e:\n            print(e, flush=True)\n\n\n## HANDEL ==\"TRAIN\":\n\ntrain_selfies = []\ntrain_image_id =[]\nwith open(selfies_train_dir) as file:\n    lines = file.readlines()\n    for line in tqdm(lines):\n        tokens = line.split(\",\")\n        image_id = str(tokens[0]) \n\n        try:\n            Selfies_1 = tokens[1].split('[')\n            Selfies=['['+x.strip() for x in Selfies_1 if x.strip() != '']\n            Selfies= \" \".join(str(i) for i in Selfies)\n            SMILES_ = \"<start> \" + Selfies + \" <end>\"\n            train_selfies.append(SMILES_)\n            train_image_id.append(image_id)\n            \n        except IndexError as e:\n            print(e, flush=True)\n    \n\n# token\nall_selfies = test_selfies +train_selfies\ntop_k = 500\ntokenizer = tf.keras.preprocessing.text.Tokenizer(\n    num_words=top_k, oov_token=\"<unk>\", filters='!\"$&:;?^`{}~ ', lower=False\n)\ntokenizer.fit_on_texts(all_selfies)\n\n\ntrain_seqs = tokenizer.texts_to_sequences(train_selfies)\ntest_seqs = tokenizer.texts_to_sequences(test_selfies)\n\ntokenizer.word_index[\"<pad>\"] = 0\ntokenizer.index_word[0] = \"<pad>\"\n\n# padding each vector to the max_length of the captions, if the max_length parameter is not provided, pad_sequences calculates that automatically\ntest_cap_vector = tf.keras.preprocessing.sequence.pad_sequences(\n    test_seqs, padding=\"post\"\n)\ntrain_cap_vector = tf.keras.preprocessing.sequence.pad_sequences(\n    train_seqs, padding=\"post\"\n)\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# train_patterns = tf.cast(train_df[\"patterns\"].values.tolist(), tf.uint8)\n# val_patterns = tf.cast(val_df[\"patterns\"].values.tolist(), tf.uint8)\n\ndisplay(train_cap_vector)\nlen(train_cap_vector[100000])","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 style=\"font-family: Verdana; font-size: 20px; font-style: normal; font-weight: normal; text-decoration: none; text-transform: none; letter-spacing: 2px; color: navy; background-color: #ffffff;\">4.3  UTILITY FUNCTIONS FOR TFRECORD CREATION</h3>","metadata":{}},{"cell_type":"code","source":"# def _bytes_feature(value, is_list=False):\n#     \"\"\"Returns a bytes_list from a string / byte.\"\"\"\n#     if isinstance(value, type(tf.constant(0))):\n#         value = value.numpy() # BytesList won't unpack a string from an EagerTensor.\n    \n#     if not is_list:\n#         value = [value]\n    \n#     return tf.train.Feature(bytes_list=tf.train.BytesList(value=value))\n\n# def _float_feature(value, is_list=False):\n#     \"\"\"Returns a float_list from a float / double.\"\"\"\n        \n#     if not is_list:\n#         value = [value]\n        \n#     return tf.train.Feature(float_list=tf.train.FloatList(value=value))\n\n# def _int64_feature(value, is_list=False):\n#     \"\"\"Returns an int64_list from a bool / enum / int / uint.\"\"\"\n        \n#     if not is_list:\n#         value = [value]\n        \n#     return tf.train.Feature(int64_list=tf.train.Int64List(value=value))\n\n\n# def create_tf_dataset(df, is_test=False, tokenized_smile=None, patterns=None):\n#     ds = tf.data.Dataset.from_tensor_slices(df[\"img_path\"].values)\n#     ds_roi_bbox = tf.data.Dataset.from_tensor_slices([tf.constant(x) for x in tqdm(df[\"roi_bbox\"].values)])\n#     ds = tf.data.Dataset.zip((ds, ds_roi_bbox))\n    \n#     if is_test:\n#         return ds\n#     else:\n#         if tokenized_smile is not None:\n#             target_ds = tf.data.Dataset.from_tensor_slices(tokenized_smile)\n#         else:\n#             target_ds = tf.data.Dataset.from_tensor_slices(df[\"smile\"].values)\n\n#         pattern_ds = tf.data.Dataset.from_tensor_slices(patterns)\n#         return tf.data.Dataset.zip((ds, target_ds, pattern_ds))\n\n    \n# def prep_tf_dataset_w_target(img_path, bbox, target, pattern, img_shape,\n#                              invert=True, \n#                              rotate_trick=True, \n#                              repair_images=False):\n    \n#     img_tensor = tf_load_image(img_path, bbox,\n#                                img_size=img_shape[:2], \n#                                tile_to_3_channel=img_shape[-1]==3, \n#                                invert=invert, \n#                                rotate_trick=rotate_trick, \n#                                repair_images=repair_images)\n#     return img_tensor, target, pattern\n\n\n# def prep_tf_dataset_wo_target(img_path, bbox, img_shape,\n#                              invert=True, \n#                              rotate_trick=True, \n#                              repair_images=False):\n    \n#     img_tensor = tf_load_image(img_path, bbox,\n#                                img_size=img_shape[:2], \n#                                tile_to_3_channel=img_shape[-1]==3, \n#                                invert=invert, \n#                                rotate_trick=rotate_trick, \n#                                repair_images=repair_images)\n#     img_id = tf.strings.split(tf.strings.split(img_path, \".png\")[0], \"/\")[-1]\n#     return img_tensor, img_id\n\n\n# def serialize_tokenized(image, smile, pattern, image_id,is_test=False):\n#     \"\"\"\n#     Creates a tf.Example message ready to be written to a file from 4 features.\n\n#     Args:\n#         image (TBD): TBD\n#         other: Either the image_id or the target smile\n    \n#     Returns:\n#         A tf.Example Message ready to be written to file\n#     \"\"\"\n#     # Create a dictionary mapping the feature name to the \n#     # tf.Example-compatible data type.\n#     const = {0:tf.constant(0, dtype=tf.uint8),1:tf.constant(1, dtype=tf.uint8)}\n#     feature = {'image': _bytes_feature(tf.io.encode_png(image), is_list=False)}\n#     if not is_test:\n#         feature[\"smile\"] = _int64_feature(smile, is_list=True)\n#         feature[\"patterns\"] = _int64_feature(pattern.tobytes(), is_list=True)\n        \n#         feature[\"image_id\"] = _bytes_feature(image_id, is_list=False)\n#     else:\n#         feature[\"image_id\"] = _bytes_feature(image_id, is_list=False)\n    \n#     # Create a Features message using tf.train.Example.\n#     example_proto = tf.train.Example(features=tf.train.Features(feature=feature))\n#     print(feature)\n#     return example_proto.SerializeToString()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# train_ds = create_tf_dataset(train_df, tokenized_smile=train_captions, patterns=train_patterns)\n# train_ds = train_ds.map(lambda x,y,z: (prep_tf_dataset_w_target(x[0], x[1], y, z, IMG_SHAPE, \n#                                                               invert=DO_INVERT, \n#                                                               rotate_trick=FIX_ROTATION, \n#                                                               repair_images=DO_REPAIR)), num_parallel_calls=tf.data.AUTOTUNE).prefetch(tf.data.AUTOTUNE)\n\n# val_ds = create_tf_dataset(val_df, tokenized_smile=val_captions, patterns=val_patterns)\n# val_ds = val_ds.map(lambda x,y,z: (prep_tf_dataset_w_target(x[0], x[1], y, z, IMG_SHAPE, \n#                                                           invert=DO_INVERT, \n#                                                           rotate_trick=FIX_ROTATION, \n#                                                           repair_images=DO_REPAIR)), num_parallel_calls=tf.data.AUTOTUNE).prefetch(tf.data.AUTOTUNE)\n\n# test_ds = create_tf_dataset(ss_df, is_test=True)\n# test_ds = test_ds.map(lambda x,y: (prep_tf_dataset_wo_target(x, y, IMG_SHAPE, \n#                                                            invert=DO_INVERT, \n#                                                            rotate_trick=FIX_ROTATION, \n#                                                            repair_images=DO_REPAIR)),num_parallel_calls=tf.data.AUTOTUNE).prefetch(tf.data.AUTOTUNE)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# def write_tfrecords(ds, n_ex, n_ex_per_rec=20000, serialize_fn=serialize_tokenized, out_dir=\"/kaggle/working/train_records\", is_test=False):\n#     n_recs = int(np.ceil(n_ex/n_ex_per_rec))\n    \n#     # Make dataset iterable\n#     ds = ds.as_numpy_iterator()\n    \n#     # Create folder\n#     if not os.path.isdir(out_dir):\n#         os.makedirs(out_dir, exist_ok=True)\n        \n#     # Create tfrecords\n#     for i in tqdm(range(n_recs), total=n_recs):\n#         print(f\"\\n... Writing TFRecord {i+1} of {n_recs} ...\\n\")\n#         tfrec_path = os.path.join(out_dir, f\"{out_dir.rsplit('_', 1)[1]}_{(i+1):02}_{n_recs:02}.tfrec\")\n#         with tf.io.TFRecordWriter(tfrec_path) as writer:\n#             for ex in tqdm(range(n_ex_per_rec), total=n_ex_per_rec):\n#                 try:\n#                     example = serialize_fn(*next(ds), is_test=is_test)\n#                     writer.write(example)\n#                 except:\n#                     break\n                    \n# print(\"\\n... MAKING TRAINING TFRECORDS ...\\n\")\n# write_tfrecords(train_ds, N_TRAIN, N_EX_PER_REC, serialize_fn=serialize_tokenized, out_dir=\"/kaggle/working/train_records\")\n                    \n# print(\"\\n... MAKING VALIDATION TFRECORDS ...\\n\")\n# write_tfrecords(val_ds, N_VAL, N_EX_PER_REC, serialize_fn=serialize_tokenized, out_dir=\"/kaggle/working/val_records\")\n\n# print(\"\\n... MAKING TESTING TFRECORDS ...\\n\")\n# write_tfrecords(test_ds, N_TEST, N_EX_PER_REC, serialize_fn=serialize_tokenized, out_dir=\"/kaggle/working/test_records\", is_test=True)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# def decode_records(serialized_example, is_test=False, is_tokenized=True, img_shape=(192,384,3)):\n#     \"\"\" Parses a set of features and label from the given `serialized_example`.\n        \n#         It is used as a map function for `dataset.map`\n\n#     Args:\n#         serialized_example (tf.Example): A serialized example containing the\n#             following features:\n#                 – sensor_feature_0 – [int64]\n#                 – sensor_feature_1 – [int64]\n#                 – sensor_feature_2 – [int64]\n#         is_test (bool, optional): Whether to allow for the label feature\n        \n#     Returns:\n#         A decoded tf.data.Dataset object representing the tfrecord dataset\n#     \"\"\"\n#     # Defaults are not specified since both keys are required.\n#     feature_dict = {\n#         'image': tf.io.FixedLenFeature(shape=[], dtype=tf.string),\n#     }\n    \n#     if not is_test:\n#         if is_tokenized:\n#             feature_dict[\"inchi\"] = tf.io.FixedLenFeature(shape=[MAX_LEN], dtype=tf.int64, default_value=[0]*MAX_LEN)\n#             feature_dict[\"patterns\"] = tf.io.FixedLenFeature(shape=[591], dtype=tf.int64, default_value=[0]*591)\n#         else:\n#             feature_dict[\"inchi\"] = tf.io.FixedLenFeature(shape=[], dtype=tf.string, default_value='')\n#     else:\n#         feature_dict['image_id'] = tf.io.FixedLenFeature(shape=[], dtype=tf.string, default_value='')\n    \n  \n#     Define a parser\n#     features = tf.io.parse_single_example(serialized_example, features=feature_dict)\n    \n#     image = decode_tf_ex_image(features['image'], resize_to=(*img_shape[:2], 3))\n#     if not is_test:\n#         inchi = features[\"inchi\"]\n#         pattern = features[\"patterns\"]\n#         return image, inchi, pattern\n#     else:\n#         image_id = features[\"image_id\"]\n#         return image, image_id\n    \n# def decode_tf_ex_image(image_data, resize_to=(192,384,3)):\n#     image = tf.image.decode_png(image_data, channels=3)\n#     image = tf.reshape(image, resize_to)\n#     return tf.cast(image, tf.uint8)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# print(\"\\n... TRAIN CHECK ...\\n\")\n# CHECK_TRAIN_TFREC_PATHS = sorted(tf.io.gfile.glob(f'/kaggle/working/train_records/*.tfrec'), key=lambda x: int(x[:-4].rsplit(\"_\", 2)[1]))[:1]\n# check_train_ds = tf.data.TFRecordDataset(CHECK_TRAIN_TFREC_PATHS, num_parallel_reads=tf.data.AUTOTUNE)\n# check_train_ds = check_train_ds.map(lambda x: (decode_records(x, img_shape=IMG_SHAPE)), num_parallel_calls=tf.data.AUTOTUNE)\n# for x,y,z in check_train_ds.take(1):\n#     print(z)\n# #     plt.figure(figsize=(12,12))\n# #     plt.imshow(x)\n# #     plt.title(\"\".join([INT2TOK[c] for c in  y.numpy() if c not in [0,1,2]]))\n# #     plt.show()\n\n    \n# print(\"\\n... VALIDATION CHECK ...\\n\")\n# CHECK_VAL_TFREC_PATHS = sorted(tf.io.gfile.glob(f'/kaggle/working/val_records/*.tfrec'), key=lambda x: int(x[:-4].rsplit(\"_\", 2)[1]))[:1]\n# check_val_ds = tf.data.TFRecordDataset(CHECK_VAL_TFREC_PATHS, num_parallel_reads=tf.data.AUTOTUNE)\n# check_val_ds = check_val_ds.map(lambda x: (decode_records(x, img_shape=IMG_SHAPE)), num_parallel_calls=tf.data.AUTOTUNE)\n# for x,y,z in check_val_ds.take(1):\n#     print(z)\n# #     plt.figure(figsize=(12,12))\n# #     plt.imshow(x)\n# #     plt.title(\"\".join([INT2TOK[c] for c in  y.numpy() if c not in [0,1,2]]))\n# #     plt.show()\n\n\n# # print(\"\\n... TEST CHECK ...\\n\")\n# # CHECK_TEST_TFREC_PATHS = sorted(tf.io.gfile.glob(f'/kaggle/working/test_records/*.tfrec'), key=lambda x: int(x[:-4].rsplit(\"_\", 2)[1]))[:1]\n# # check_test_ds = tf.data.TFRecordDataset(CHECK_TEST_TFREC_PATHS, num_parallel_reads=tf.data.AUTOTUNE)\n# # check_test_ds = check_test_ds.map(lambda x: (decode_records(x, img_shape=IMG_SHAPE, is_test=True)), num_parallel_calls=tf.data.AUTOTUNE)\n# # for x,y in check_test_ds.take(1):\n# #     plt.figure(figsize=(12,12))\n# #     plt.imshow(x)\n# #     plt.title(str(y.numpy().decode()))\n# #     plt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}