{"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":"**Original Png Dataset: [RSNA Breast Cancer Detection - 1024x1024 pngs](https://www.kaggle.com/datasets/theoviel/rsna-breast-cancer-1024-pngs)**\n<div class='alert alert-block alert-success'>\n    <h4>Original Images DataSets:</h4>\n    <ul>\n        <li><a href='https://www.kaggle.com/datasets/olegbaryshnikov/rsna-resized-tfrecords-1024x1024'><b>1024x1024 TFRecords</b></a></li>\n        <li><a href='https://www.kaggle.com/datasets/olegbaryshnikov/rsna-resized-tfrecords-768x768'><b>768x768 TFRecords</b></a></li>\n        <li><a href='https://www.kaggle.com/datasets/olegbaryshnikov/rsna-resized-tfrecords-512x512'><b>512x512 TFRecords</b></a></li>\n    </ul>\n</div>\n\n**ROI Dataset: [[RSNA] ROI 1024x1024 pngs](https://www.kaggle.com/datasets/olegbaryshnikov/rsna-roi-1024x1024-pngs) extracted using [[RSNA] Improved ROI Extraction [YOLOv5]](https://www.kaggle.com/code/olegbaryshnikov/rsna-improved-roi-extraction-yolov5) notebook**\n<div class='alert alert-block alert-success'>\n    <h4>1:1 ROI DataSets:</h4>\n    <ul>\n        <li><a href='https://www.kaggle.com/datasets/olegbaryshnikov/rsna-roi-tfrecords-1024x1024'><b>ROI 1024x1024 TFRecords</b></a></li>\n        <li><a href='https://www.kaggle.com/datasets/olegbaryshnikov/rsna-roi-tfrecords-768x768'><b>ROI 768x768 TFRecords</b></a></li>\n        <li><a href='https://www.kaggle.com/datasets/olegbaryshnikov/rsna-roi-tfrecords-512x512'><b>ROI 512x512 TFRecords</b></a></li>\n    </ul>\n</div>\n\n**ROI Dataset: [[RSNA] ROI 512x1024 pngs](https://www.kaggle.com/datasets/olegbaryshnikov/rsna-roi-512x1024-pngs) extracted using [[RSNA] Improved ROI Extraction [YOLOv5]](https://www.kaggle.com/code/olegbaryshnikov/rsna-improved-roi-extraction-yolov5) notebook**\n<div class='alert alert-block alert-success'>\n    <h4>1:2 ROI DataSets <span style=\"color: orange\">[Recommended]</span>:</h4>\n    <ul>\n        <li><a href='https://www.kaggle.com/datasets/olegbaryshnikov/rsna-roi-tfrecords-512x1024'><b>ROI 512x1024 TFRecords</b></a></li>\n        <li><a href='https://www.kaggle.com/datasets/olegbaryshnikov/rsna-roi-tfrecords-384x768'><b>ROI 384x768 TFRecords</b></a></li>\n        <li><a href='https://www.kaggle.com/datasets/olegbaryshnikov/rsna-roi-tfrecords-256x512'><b>ROI 256x512 TFRecords</b></a></li>\n    </ul>\n</div>\n\n<div class='alert alert-block alert-danger' style='overflow: auto;'>\n    <b>\n        <div style='float:left;height:100%;width:5%'>⛔</div>\n        <div style='float:left;height:100%;width:90%'>Don't forget to use 'GZIP' compression_type option, when using TFRecordDataset. Check out the last section of the notebook for usage examples.</div>\n        <div style='float:left;height:100%;width:5%'>⛔</div>\n    </b>\n</div>\n\n<div class='alert alert-block alert-info'>\n    <b>\n        🔷 Some data to save changes during preprocessing! 🔷\n        <ul>\n            <li>the 'laterality' and 'view' columns are label encoded, check out label encoder's vocabularies below.</li>\n            <li>nan values for the 'age' column (for machine_id == 49) are changed to the median value.</li>\n        </ul>\n    </b>\n</div>\n\n**TFRecord data is already split using StratifiedKFold. Use KFold on generated TFRecords for cross-validation.**\n\n<div class='alert alert-block alert-info'>\n    <b>\n        <ul>\n            <div>Original Data:</div>\n            <li>\n                512x512 Original Data TFRecords:\n                <a href='https://www.kaggle.com/code/olegbaryshnikov/rsna-tfrecods-resized-images-and-roi?scriptVersionId=113228702'>[Notebook version]</a>\n            </li>\n            <li>\n                768x768 Original Data TFRecords:\n                <a href='https://www.kaggle.com/code/olegbaryshnikov/rsna-tfrecods-resized-images-and-roi?scriptVersionId=113228720'>[Notebook version]</a>\n            </li>\n            <p>\n            <div>1:1 ROI Data:</div>\n            <li>\n                512x512 ROI TFRecords:\n                <a href='https://www.kaggle.com/code/olegbaryshnikov/rsna-tfrecods-resized-images-and-roi?scriptVersionId=113577136'>[Notebook version]</a>\n            </li>\n            <li>\n                768x768 ROI TFRecords:\n                <a href='https://www.kaggle.com/code/olegbaryshnikov/rsna-tfrecods-resized-images-and-roi?scriptVersionId=113577147'>[Notebook version]</a>\n            </li>\n            <p>\n            <div>1:2 ROI Data:</div>\n            <li>\n                384x768 ROI TFRecords:\n                <a href='https://www.kaggle.com/code/olegbaryshnikov/rsna-tfrecods-resized-images-and-roi?scriptVersionId=116867931'>[Notebook version]</a>\n            </li>\n            <li>\n                256x512 ROI TFRecords:\n                <a href='https://www.kaggle.com/code/olegbaryshnikov/rsna-tfrecods-resized-images-and-roi?scriptVersionId=116867977'>[Notebook version]</a>\n            </li>\n        </ul>\n    </b>\n</div>","metadata":{}},{"cell_type":"markdown","source":"## Imports","metadata":{}},{"cell_type":"code","source":"import os\nimport math\nimport random\nimport numpy as np\nimport pandas as pd\nfrom IPython.display import display\n\nimport tensorflow as tf\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\nimport cv2\n\nfrom tqdm.notebook import tqdm\nimport gc\n\nimport glob\nfrom sklearn.preprocessing import LabelEncoder\n\nfrom joblib import Parallel, delayed\nfrom sklearn.model_selection import StratifiedGroupKFold","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-01-20T12:45:10.936435Z","iopub.execute_input":"2023-01-20T12:45:10.937223Z","iopub.status.idle":"2023-01-20T12:45:10.944656Z","shell.execute_reply.started":"2023-01-20T12:45:10.93717Z","shell.execute_reply":"2023-01-20T12:45:10.94337Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Config","metadata":{}},{"cell_type":"code","source":"Config = {\n    'output_dim_x' : 128,\n    'output_dim_y' : 256,\n    'max_fold_num' : 120,\n    'seed': 1111,\n}","metadata":{"execution":{"iopub.status.busy":"2023-01-20T12:45:10.946652Z","iopub.execute_input":"2023-01-20T12:45:10.947058Z","iopub.status.idle":"2023-01-20T12:45:10.955856Z","shell.execute_reply.started":"2023-01-20T12:45:10.947021Z","shell.execute_reply":"2023-01-20T12:45:10.954613Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def set_seed(seed=42):\n    np.random.seed(seed)\n    random.seed(seed)\n    tf.random.set_seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    \nset_seed(Config[\"seed\"])","metadata":{"execution":{"iopub.status.busy":"2023-01-20T12:45:10.957268Z","iopub.execute_input":"2023-01-20T12:45:10.958083Z","iopub.status.idle":"2023-01-20T12:45:10.968013Z","shell.execute_reply.started":"2023-01-20T12:45:10.958045Z","shell.execute_reply":"2023-01-20T12:45:10.966472Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_csv_path = '/kaggle/input/rsna-breast-cancer-detection/train.csv'\ntrain_images_folder = '/kaggle/input/rsna-roi-512x1024-pngs'\n\nunrec_img_csv_path = '/kaggle/input/rsna-roi-512x1024-pngs/unrecognized_images.csv'","metadata":{"execution":{"iopub.status.busy":"2023-01-20T12:45:10.970459Z","iopub.execute_input":"2023-01-20T12:45:10.971592Z","iopub.status.idle":"2023-01-20T12:45:10.979077Z","shell.execute_reply.started":"2023-01-20T12:45:10.971549Z","shell.execute_reply":"2023-01-20T12:45:10.977724Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data preprocessing","metadata":{}},{"cell_type":"code","source":"train_csv=pd.read_csv(train_csv_path)\ntrain_csv","metadata":{"execution":{"iopub.status.busy":"2023-01-20T12:45:10.980964Z","iopub.execute_input":"2023-01-20T12:45:10.981365Z","iopub.status.idle":"2023-01-20T12:45:11.149617Z","shell.execute_reply.started":"2023-01-20T12:45:10.981318Z","shell.execute_reply":"2023-01-20T12:45:11.148152Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def label_encode(col_name):\n    le = LabelEncoder()\n    train_csv[col_name]=le.fit_transform(train_csv[col_name])\n\n    le_dict = dict(zip(le.classes_, le.transform(le.classes_)))\n    display(le_dict)\n\n    display(train_csv)","metadata":{"execution":{"iopub.status.busy":"2023-01-20T12:45:11.151663Z","iopub.execute_input":"2023-01-20T12:45:11.152086Z","iopub.status.idle":"2023-01-20T12:45:11.158556Z","shell.execute_reply.started":"2023-01-20T12:45:11.152049Z","shell.execute_reply":"2023-01-20T12:45:11.157056Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#laterality and view label encoding\nto_le = ['laterality','view']\n\nfor col_name in to_le:\n    label_encode(col_name)","metadata":{"execution":{"iopub.status.busy":"2023-01-20T12:45:11.160439Z","iopub.execute_input":"2023-01-20T12:45:11.160848Z","iopub.status.idle":"2023-01-20T12:45:11.249652Z","shell.execute_reply.started":"2023-01-20T12:45:11.160802Z","shell.execute_reply":"2023-01-20T12:45:11.248376Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#BIRADS and density label encoding to remove NaN\nto_le = ['BIRADS','density']\n\nfor col_name in to_le:\n    label_encode(col_name)","metadata":{"execution":{"iopub.status.busy":"2023-01-20T12:45:11.25161Z","iopub.execute_input":"2023-01-20T12:45:11.252012Z","iopub.status.idle":"2023-01-20T12:45:11.319916Z","shell.execute_reply.started":"2023-01-20T12:45:11.251976Z","shell.execute_reply":"2023-01-20T12:45:11.318679Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#some ages (for machine_id==49) has nan values\nprint('Number of nan values: ',len(train_csv[train_csv['age'].isnull()]['age']))\nprint('Number of nan values for machine_id 49: ',len(train_csv[(train_csv['machine_id']==49)&(train_csv['age'].isnull())]['age']))\nprint('Mean age: ',np.mean(train_csv['age']))\nprint('Median age: ',np.nanmedian(train_csv['age']))\nprint('Mean age for machine_id 49: ',np.mean(train_csv[train_csv['machine_id']==49]['age']))\nprint('Median age for machine_id 49: ',np.nanmedian(train_csv[train_csv['machine_id']==49]['age']))","metadata":{"execution":{"iopub.status.busy":"2023-01-20T12:45:11.32157Z","iopub.execute_input":"2023-01-20T12:45:11.321994Z","iopub.status.idle":"2023-01-20T12:45:11.354462Z","shell.execute_reply.started":"2023-01-20T12:45:11.321957Z","shell.execute_reply":"2023-01-20T12:45:11.353406Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#assign median age for machine_id 49 as age for age nan values\ntrain_csv.loc[train_csv['age'].isnull(),'age'] = np.nanmedian(train_csv[train_csv['machine_id']==49]['age'])\ntrain_csv['age'].isnull().values.any()","metadata":{"execution":{"iopub.status.busy":"2023-01-20T12:45:11.358547Z","iopub.execute_input":"2023-01-20T12:45:11.358947Z","iopub.status.idle":"2023-01-20T12:45:11.376579Z","shell.execute_reply.started":"2023-01-20T12:45:11.358913Z","shell.execute_reply":"2023-01-20T12:45:11.375124Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def show_image(img):\n    fig=plt.figure(figsize=(10, 10))\n    plt.imshow(img, cmap='bone')\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-01-20T12:45:11.378472Z","iopub.execute_input":"2023-01-20T12:45:11.37896Z","iopub.status.idle":"2023-01-20T12:45:11.389436Z","shell.execute_reply.started":"2023-01-20T12:45:11.378913Z","shell.execute_reply":"2023-01-20T12:45:11.386201Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def read_png_img(patient_id, image_id):\n    img_path = os.path.join(train_images_folder,f'{patient_id}_{image_id}.png')\n    img = cv2.imread(img_path)\n    \n    return img","metadata":{"execution":{"iopub.status.busy":"2023-01-20T12:45:11.390851Z","iopub.execute_input":"2023-01-20T12:45:11.391278Z","iopub.status.idle":"2023-01-20T12:45:11.398518Z","shell.execute_reply.started":"2023-01-20T12:45:11.391239Z","shell.execute_reply":"2023-01-20T12:45:11.39724Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"unrec_img = pd.read_csv(unrec_img_csv_path, index_col=0)\nunrec_img","metadata":{"execution":{"iopub.status.busy":"2023-01-20T12:45:11.40023Z","iopub.execute_input":"2023-01-20T12:45:11.400643Z","iopub.status.idle":"2023-01-20T12:45:11.429676Z","shell.execute_reply.started":"2023-01-20T12:45:11.400609Z","shell.execute_reply":"2023-01-20T12:45:11.428453Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#drop unrecognized:\ntrain_csv_index = train_csv.set_index(['patient_id','image_id']).index\nunrec_img_index = unrec_img.set_index(['patient_id','image_id']).index\nmask = ~train_csv_index.isin(unrec_img_index)\ntrain_csv = train_csv[mask].reset_index(drop=True)\ntrain_csv","metadata":{"execution":{"iopub.status.busy":"2023-01-20T12:45:11.431663Z","iopub.execute_input":"2023-01-20T12:45:11.432118Z","iopub.status.idle":"2023-01-20T12:45:11.511588Z","shell.execute_reply.started":"2023-01-20T12:45:11.432076Z","shell.execute_reply":"2023-01-20T12:45:11.51049Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Convert to TFRecords","metadata":{}},{"cell_type":"code","source":"split_col = 'cancer'\ngroup_col = 'patient_id'\n\n\nskf = StratifiedGroupKFold(n_splits=Config['max_fold_num'], random_state=Config[\"seed\"], shuffle=True)\n\nfor fold, (_, val_ind) in enumerate(skf.split(X=train_csv, y=train_csv[split_col], groups=train_csv[group_col])):\n    train_csv.loc[val_ind,'split_num'] = fold\n    \ntrain_csv['split_num'] = np.int32(train_csv['split_num'])\ntrain_csv","metadata":{"execution":{"iopub.status.busy":"2023-01-20T12:47:56.076953Z","iopub.execute_input":"2023-01-20T12:47:56.077828Z","iopub.status.idle":"2023-01-20T12:50:31.630848Z","shell.execute_reply.started":"2023-01-20T12:47:56.077777Z","shell.execute_reply":"2023-01-20T12:50:31.629551Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def _bytes_feature(value):\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  return tf.train.Feature(bytes_list=tf.train.BytesList(value=[value]))\n\ndef _float_feature(value):\n  \"\"\"Returns a float_list from a float / double.\"\"\"\n  return tf.train.Feature(float_list=tf.train.FloatList(value=[value]))\n\ndef _int64_feature(value):\n  \"\"\"Returns an int64_list from a bool / enum / int / uint.\"\"\"\n  return tf.train.Feature(int64_list=tf.train.Int64List(value=[value]))","metadata":{"execution":{"iopub.status.busy":"2023-01-20T12:47:43.599117Z","iopub.execute_input":"2023-01-20T12:47:43.599564Z","iopub.status.idle":"2023-01-20T12:47:43.607959Z","shell.execute_reply.started":"2023-01-20T12:47:43.599528Z","shell.execute_reply":"2023-01-20T12:47:43.606655Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def serialize_example(image, cancer, laterality, view, age, implant, machine_id, site_id):\n  # Create a dictionary mapping the feature name to the tf.train.Example-compatible\n  # data type.\n  feature = {\n      'image': _bytes_feature(image),\n      'cancer': _int64_feature(cancer),\n      'laterality': _int64_feature(laterality),\n      'view': _int64_feature(view),\n      'age': _float_feature(age),\n      'implant': _int64_feature(implant),\n      'machine_id': _int64_feature(machine_id),\n      'site_id': _int64_feature(site_id)\n  }\n\n  # Create a Features message using tf.train.Example.\n\n  example_proto = tf.train.Example(features=tf.train.Features(feature=feature))\n  return example_proto.SerializeToString()","metadata":{"execution":{"iopub.status.busy":"2023-01-20T12:47:43.60956Z","iopub.execute_input":"2023-01-20T12:47:43.609898Z","iopub.status.idle":"2023-01-20T12:47:43.624527Z","shell.execute_reply.started":"2023-01-20T12:47:43.609868Z","shell.execute_reply":"2023-01-20T12:47:43.62335Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def process(split_num):\n    dim_x = Config['output_dim_x']\n    dim_y = Config['output_dim_y']\n    col_names = ['patient_id', 'image_id', 'cancer', 'laterality', 'view', 'age', 'implant', 'machine_id', 'site_id']\n    \n    train_split = train_csv[train_csv['split_num']==split_num][col_names]\n    train_split = train_split.sample(frac=1)\n    \n    filename = f'train_{dim_x}x{dim_y}/split_{split_num}.tfrec'\n    options = tf.io.TFRecordOptions(compression_type = 'GZIP')\n\n    with tf.io.TFRecordWriter(filename, options) as writer:\n        for data in train_split.itertuples():\n            img = read_png_img(str(data.patient_id),str(data.image_id))\n\n            #show_image(img)\n            #resizing\n            img_resized = cv2.resize(img, (dim_x, dim_y),interpolation = cv2.INTER_AREA)\n            #img_resized = cv2.imencode('.png', img_resized, (cv2.IMWRITE_PNG_COMPRESSION , 9))[1]\n            #img_resized = tf.io.encode_png(img_resized, compression = 9).numpy()\n            #show_image(img_resized)\n\n            #convert to tfrec\n            example = serialize_example(img_resized.tobytes(), data.cancer, data.laterality,\n                                        data.view, data.age, data.implant, data.machine_id, data.site_id)\n            writer.write(example)\n    writer.close()\n    gc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-01-20T12:47:43.626575Z","iopub.execute_input":"2023-01-20T12:47:43.62705Z","iopub.status.idle":"2023-01-20T12:47:43.638646Z","shell.execute_reply.started":"2023-01-20T12:47:43.627004Z","shell.execute_reply":"2023-01-20T12:47:43.637108Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#check png files\nfig=plt.figure(figsize=(25, 5))\n\nfor i,data in enumerate(train_csv.iloc[195:200].itertuples()):\n    img = read_png_img(str(data.patient_id),str(data.image_id))\n    #display(img)\n    fig.add_subplot(1, 5, i+1)\n    \n    plt.imshow(img, cmap='bone')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-01-20T12:47:43.640607Z","iopub.execute_input":"2023-01-20T12:47:43.641445Z","iopub.status.idle":"2023-01-20T12:47:44.728858Z","shell.execute_reply.started":"2023-01-20T12:47:43.641393Z","shell.execute_reply":"2023-01-20T12:47:44.727562Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.makedirs(f'train_{Config[\"output_dim_x\"]}x{Config[\"output_dim_y\"]}', exist_ok = True)\n\n_ = Parallel(n_jobs=-1)(\n    delayed(process)(split_num)\n    for split_num in tqdm(range(0,Config['max_fold_num']))\n)","metadata":{"execution":{"iopub.status.busy":"2023-01-20T12:47:44.730362Z","iopub.execute_input":"2023-01-20T12:47:44.730746Z","iopub.status.idle":"2023-01-20T12:47:49.248199Z","shell.execute_reply.started":"2023-01-20T12:47:44.73071Z","shell.execute_reply":"2023-01-20T12:47:49.246253Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Check TFRecords","metadata":{}},{"cell_type":"code","source":"os.system(f'printf \"Number of files in train_{Config[\"output_dim_x\"]}x{Config[\"output_dim_y\"]}: $(ls /kaggle/working/train_{Config[\"output_dim_x\"]}x{Config[\"output_dim_y\"]} | wc -l)\\n\"')","metadata":{"execution":{"iopub.status.busy":"2023-01-20T12:47:49.249816Z","iopub.status.idle":"2023-01-20T12:47:49.250315Z","shell.execute_reply.started":"2023-01-20T12:47:49.250078Z","shell.execute_reply":"2023-01-20T12:47:49.250101Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def _parse_image_function(example_proto, dim_x, dim_y):\n    image_feature_description = {\n        'image': tf.io.FixedLenFeature([], tf.string),\n        'cancer': tf.io.FixedLenFeature([], tf.int64),\n        'laterality': tf.io.FixedLenFeature([], tf.int64),\n        'view': tf.io.FixedLenFeature([], tf.int64),\n        'age': tf.io.FixedLenFeature([], tf.float32),\n        'implant': tf.io.FixedLenFeature([], tf.int64),\n        'machine_id': tf.io.FixedLenFeature([], tf.int64),\n        'site_id': tf.io.FixedLenFeature([], tf.int64),\n    }\n    \n    single_example = tf.io.parse_single_example(example_proto, image_feature_description)\n    \n    #for png images\n    #image = tf.reshape(tf.io.decode_png(single_example['image'],dtype=np.dtype('uint8')), (dim,dim,3))\n    #for raw images\n    image = tf.reshape(tf.io.decode_raw(single_example['image'],out_type=np.dtype('uint8')), (dim_y,dim_x,3))\n    cancer =  single_example['cancer']\n    laterality =  single_example['laterality']\n    view =  single_example['view']\n    age =  single_example['age']\n    implant =  single_example['implant']\n    machine_id =  single_example['machine_id']\n    site_id =  single_example['site_id']\n    \n    inputs = (image,laterality,view,age,implant,machine_id,site_id)\n    labels = cancer\n    \n    return ((inputs),labels)\n\n\ndef load_dataset(filenames,dim_x,dim_y):\n    dataset = tf.data.TFRecordDataset(filenames, compression_type = 'GZIP')\n    dataset = dataset.map(lambda ex: _parse_image_function(ex,dim_x,dim_y))\n    return dataset\n\nBATCH_NUM=8\ndef get_dataset(FILENAME,dim_x,dim_y):\n    dataset = load_dataset(FILENAME,dim_x,dim_y)\n    dataset = dataset.shuffle(64).batch(BATCH_NUM)\n    return dataset","metadata":{"execution":{"iopub.status.busy":"2023-01-20T12:47:49.25337Z","iopub.status.idle":"2023-01-20T12:47:49.254232Z","shell.execute_reply.started":"2023-01-20T12:47:49.253914Z","shell.execute_reply":"2023-01-20T12:47:49.253945Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_tfrecords = glob.glob(f'train_{Config[\"output_dim_x\"]}x{Config[\"output_dim_y\"]}/*.tfrec')\ndataset_tuple=next(iter(get_dataset(train_tfrecords,Config[\"output_dim_x\"],Config[\"output_dim_y\"])))\n\n#convert tuple to dict\ndataset_dict = {key:dataset_tuple[0][key_num] for key_num,key in enumerate(['image','laterality','view','age','implant','machine_id','site_id'])}\ndataset_dict['cancer'] = dataset_tuple[1]\n\nprint('image shape:',dataset_dict['image'].shape)\ndisplay({key:val for key, val in dataset_dict.items() if key != 'image'})\n\nfig=plt.figure(figsize=(20, 10))\nfor i in range(0,BATCH_NUM):\n    fig.add_subplot(2,BATCH_NUM//2,i+1)\n    plt.imshow(dataset_dict['image'][i], cmap='bone')\nplt.show()\n","metadata":{"execution":{"iopub.status.busy":"2023-01-20T12:47:49.2558Z","iopub.status.idle":"2023-01-20T12:47:49.256414Z","shell.execute_reply.started":"2023-01-20T12:47:49.25608Z","shell.execute_reply":"2023-01-20T12:47:49.256108Z"},"trusted":true},"execution_count":null,"outputs":[]}]}