{"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":"code","source":"# writing a dataset (16 bit) in tfrecord format for fast dataloading\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport glob\nimport os\nfrom math import floor\nimport tensorflow as tf\nfrom tensorflow.train import BytesList, FloatList, Int64List\nfrom tensorflow.train import Example, Features, Feature\n\n# ...\n# load the data\n\nBASE_DIR = '/kaggle/input/google-research-identify-contrails-reduce-global-warming/'\nTRAIN_DIR = BASE_DIR + 'train'\nTEST_DIR = BASE_DIR + 'test'\nVALIDATION_DIR = BASE_DIR + 'validation'\nOUTDIR = '/kaggle/working/train/'\nNDATASET=2\n!mkdir train\nrecord_ids= sorted(glob.glob(os.path.join(TRAIN_DIR, '*')))\n\nmean_var = [(233.68251494476192, 49.41609113053796),\n    (242.25141679731962, 84.3667759851496),\n    (250.73854949703414, 129.34094451279213),\n    (274.38815456958827, 386.2778745820613),\n    (255.53700834960705, 172.70252117076805),\n    (276.5770302843285, 430.757200225114),\n    (275.3361689617875, 447.36708595674526),\n    (272.5409457075911, 424.7027156746369),\n    (260.4165092867757, 251.74876866159116)]\n\n# number of standard deviations used for normalization\nnstd = 1\n\ndef norm_arr(arr, mean, std, num_std):\n    return (arr - mean) / (std * num_std)\n\ndef get_image(i):\n    bands = []\n    with open(os.path.join(record_ids[i], 'band_08.npy'), 'rb') as f:\n        bands.append(norm_arr(np.load(f).astype(np.float16), mean_var[0][0], mean_var[0][1]**.5, nstd))\n    with open(os.path.join( record_ids[i], 'band_09.npy'), 'rb') as f:\n        bands.append(norm_arr(np.load(f).astype(np.float16), mean_var[1][0], mean_var[1][1]**.5, nstd))\n    with open(os.path.join( record_ids[i] , 'band_10.npy'), 'rb') as f:\n        bands.append(norm_arr(np.load(f).astype(np.float16), mean_var[2][0], mean_var[2][1]**.5, nstd))\n    with open(os.path.join(record_ids[i], 'band_11.npy'), 'rb') as f:\n        bands.append(norm_arr(np.load(f).astype(np.float16), mean_var[3][0], mean_var[3][1]**.5, nstd))\n    with open(os.path.join( record_ids[i], 'band_12.npy'), 'rb') as f:\n        bands.append(norm_arr(np.load(f).astype(np.float16), mean_var[4][0], mean_var[4][1]**.5, nstd))\n    with open(os.path.join( record_ids[i] , 'band_13.npy'), 'rb') as f:\n        bands.append(norm_arr(np.load(f).astype(np.float16), mean_var[5][0], mean_var[5][1]**.5, nstd))\n    with open(os.path.join(record_ids[i], 'band_14.npy'), 'rb') as f:\n        bands.append(norm_arr(np.load(f).astype(np.float16), mean_var[6][0], mean_var[6][1]**.5, nstd))\n    with open(os.path.join( record_ids[i], 'band_15.npy'), 'rb') as f:\n        bands.append(norm_arr(np.load(f).astype(np.float16), mean_var[7][0], mean_var[7][1]**.5, nstd))\n    with open(os.path.join( record_ids[i] , 'band_16.npy'), 'rb') as f:\n        bands.append(norm_arr(np.load(f).astype(np.float16), mean_var[8][0], mean_var[8][1]**.5, nstd))\n    with open(os.path.join(record_ids[i] , 'human_pixel_masks.npy'), 'rb') as f:\n        human_pixel_mask = np.load(f).astype(bool)\n    with open(os.path.join(record_ids[i] , 'human_individual_masks.npy'), 'rb') as f:\n        human_individual_mask = np.load(f).astype(bool)\n\n    return np.stack([bands])[0], human_pixel_mask[:,:,0], human_individual_mask[:,:,0,:]\n\n# ...\n\n\ndef create_example(record_id, image_data, mask):\n    feature = {\n        'record_id': Feature(int64_list=Int64List(value=record_id)),\n        'image_data': Feature(bytes_list=BytesList(value=[image_data.tobytes()])),\n        'mask': Feature(bytes_list=BytesList(value=[mask.tobytes()])),\n    }\n    return Example(features=Features(feature=feature))\n\n# ...\ndef write_tfrecord(filename, record_ids, batch_data, masks):\n    with tf.io.TFRecordWriter(filename) as writer:\n        example = create_example(record_ids, batch_data, masks)\n        writer.write(example.SerializeToString())","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-07-10T13:05:19.970523Z","iopub.execute_input":"2023-07-10T13:05:19.971055Z","iopub.status.idle":"2023-07-10T13:05:21.250204Z","shell.execute_reply.started":"2023-07-10T13:05:19.971022Z","shell.execute_reply":"2023-07-10T13:05:21.248833Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def _parse_tfrecord(example_proto):\n    # describe how the TFRecord example will be interpreted\n    features = {'image_data': tf.io.FixedLenFeature((), tf.string),\n                'mask': tf.io.FixedLenFeature((), tf.string),\n                'record_id': tf.io.FixedLenFeature((), tf.int64)}\n    # parse the example (dict of features) from the TFRecord\n    parsed_features = tf.io.parse_single_example(example_proto, features)\n    # decode the bytes as float16 array\n    return parsed_features['record_id'], tf.io.decode_raw(parsed_features['image_data'], tf.float16), tf.io.decode_raw(parsed_features['mask'], tf.bool)\n\n\ndef parse_example(serialized_example):\n    feature_description = {\n        'record_id': tf.io.FixedLenFeature([32], tf.int64),\n        'image_data': tf.io.FixedLenFeature([], tf.string),\n        'mask': tf.io.FixedLenFeature([], tf.string),\n    }\n    example = tf.io.parse_single_example(serialized_example, feature_description)\n    \n    record_id = example['record_id']\n    image_data = tf.io.decode_raw(example['image_data'], tf.float16)\n    mask = tf.io.decode_raw(example['mask'], tf.bool)\n    \n    return record_id, image_data, mask\n\n\ndef tfrecord_input_fn(fn):\n    # read the dataset\n    dataset = tf.data.TFRecordDataset(fn)\n    # parse each example of the dataset\n    #dataset = dataset.map(_parse_tfrecord)\n    dataset = dataset.map(parse_example)\n    return dataset\n\n","metadata":{"execution":{"iopub.status.busy":"2023-07-10T13:41:11.007874Z","iopub.execute_input":"2023-07-10T13:41:11.008294Z","iopub.status.idle":"2023-07-10T13:41:11.020853Z","shell.execute_reply.started":"2023-07-10T13:41:11.008261Z","shell.execute_reply":"2023-07-10T13:41:11.019502Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n#data_np = np.array(np.random.rand(32, 256, 256, 9), dtype=np.float16)\n#ids = np.array(np.random.rand(32), dtype=np.int64)\n#masks = np.array(np.random.rand(32, 256, 256), dtype=bool)\n#write_tfrecord('test.tfrecord', ids, data_np, masks)\n\n#imgshape = (32, 256, 256, 9)\n#maskshape = (32, 256, 256)\n\n# get an iterator over the TFRecord\n#it = tfrecord_input_fn('test.tfrecord')\n#for recovered_data in it:\n#    print(np.array_equal(np.frombuffer(recovered_data[1], dtype = np.float16).reshape(imgshape) ,data_np))\n#    print(recovered_data[0] == ids)\n#    print(np.array_equal(np.frombuffer(recovered_data[2], dtype = bool).reshape(maskshape) ,masks))\n\n","metadata":{"execution":{"iopub.status.busy":"2023-07-10T13:52:40.39006Z","iopub.execute_input":"2023-07-10T13:52:40.390409Z","iopub.status.idle":"2023-07-10T13:52:40.396297Z","shell.execute_reply.started":"2023-07-10T13:52:40.390382Z","shell.execute_reply":"2023-07-10T13:52:40.395248Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ...\nmaskind = 4\nimg0, mask0, indmask0 = get_image(0)\nimg0 = np.moveaxis(img0[:,:,:,maskind], 0, -1)\nprint(img0.shape)\nprint(mask0.shape)\nprint(indmask0.shape)\n\nbatch_size = 32*8\n\ndef batch_num(ind):\n    return int(floor(float(ind)/batch_size))\n\n#num_images = len(record_ids)\nnimages = 2016*8\nstart = nimages*(NDATASET-1)\nnum_images = len(record_ids) - start\nprint(len(record_ids))\nbatch_data = np.zeros((batch_size, img0.shape[0], img0.shape[1], img0.shape[2]), dtype=img0.dtype)\nbatch_record_ids = np.zeros((batch_size,), dtype=np.int64)\nmasks = np.zeros((batch_size, mask0.shape[0], mask0.shape[1]), dtype=np.bool)\n\n# Loop over all indices in record_ids\nfor i in range(start, start+num_images):\n    if i % batch_size == 0 and i != start:\n        # Write the batch data to a TFRecord file\n        filename = OUTDIR + f\"batch_{batch_num(i-1)}.tfrecord\"\n        write_tfrecord(filename, batch_record_ids, batch_data, masks)\n\n    # Get the image data for the current hndex\n    d, m, _ = get_image(i)\n    d = np.moveaxis(d[:,:,:,maskind], 0, -1)\n    batch_data[i%batch_size] = d\n    folder_name = os.path.basename(record_ids[i])\n\n    batch_record_ids[i%batch_size] = np.int64(folder_name)\n    masks[i%batch_size] = m\n\nlast_ind = start + num_images - 1\nfilename = OUTDIR + f\"batch_{batch_num(last_ind)}.tfrecord\"\nwrite_tfrecord(filename, batch_record_ids[0:last_ind%batch_size+1], batch_data[0:last_ind%batch_size+1], masks[0:last_ind%batch_size+1])","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-07-08T14:33:41.784544Z","iopub.execute_input":"2023-07-08T14:33:41.784965Z","iopub.status.idle":"2023-07-08T14:33:43.531924Z","shell.execute_reply.started":"2023-07-08T14:33:41.784919Z","shell.execute_reply":"2023-07-08T14:33:43.53028Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-07-08T14:33:41.784544Z","iopub.execute_input":"2023-07-08T14:33:41.784965Z","iopub.status.idle":"2023-07-08T14:33:43.531924Z","shell.execute_reply.started":"2023-07-08T14:33:41.784919Z","shell.execute_reply":"2023-07-08T14:33:43.53028Z"},"trusted":true},"execution_count":null,"outputs":[]}]}