{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":6125,"sourceType":"modelInstanceVersion","modelInstanceId":4596},{"sourceId":6127,"sourceType":"modelInstanceVersion","modelInstanceId":4598},{"sourceId":6124,"sourceType":"modelInstanceVersion","modelInstanceId":4599}],"dockerImageVersionId":30732,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Setup","metadata":{}},{"cell_type":"code","source":"import os\n\nos.environ[\"KERAS_BACKEND\"] = \"jax\"  # @param [\"tensorflow\", \"jax\", \"torch\"]\n\nfrom tensorflow import data as tf_data\nimport tensorflow_datasets as tfds\nimport keras\nimport keras_cv\nimport numpy as np\nfrom keras_cv import bounding_box\nimport os\nfrom keras_cv import visualization\nimport tqdm\nimport pandas as pd\nimport pydicom\nimport tensorflow as tf\nimport tensorflow_io as tfio\n\nimport warnings\nwarnings.filterwarnings('ignore')","metadata":{"execution":{"iopub.status.busy":"2024-06-11T11:32:31.078507Z","iopub.execute_input":"2024-06-11T11:32:31.078906Z","iopub.status.idle":"2024-06-11T11:32:55.314008Z","shell.execute_reply.started":"2024-06-11T11:32:31.078867Z","shell.execute_reply":"2024-06-11T11:32:55.312864Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BASE_DIR = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/'\nTRAIN_DIR = BASE_DIR+'train_images/'\nTEST_DIR = BASE_DIR+'test_images/'\n\nPRETRAINED = 'efficientnetv2_s_imagenet'\n\nSPLIT_RATIO = .2\nBATCH_SIZE = 32\nEPOCH = 8\n\nIMG_SIZE = [320,320]","metadata":{"execution":{"iopub.status.busy":"2024-06-11T11:32:55.316197Z","iopub.execute_input":"2024-06-11T11:32:55.316863Z","iopub.status.idle":"2024-06-11T11:32:55.322868Z","shell.execute_reply.started":"2024-06-11T11:32:55.316822Z","shell.execute_reply":"2024-06-11T11:32:55.321624Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data processing","metadata":{}},{"cell_type":"code","source":"studies = os.listdir(TRAIN_DIR)\n\ntrain_studies = studies[:int(len(studies)*(1-SPLIT_RATIO))]\nval_studies = studies[int(len(studies)*(1-SPLIT_RATIO)):]","metadata":{"execution":{"iopub.status.busy":"2024-06-11T11:32:55.324409Z","iopub.execute_input":"2024-06-11T11:32:55.324753Z","iopub.status.idle":"2024-06-11T11:32:55.462075Z","shell.execute_reply.started":"2024-06-11T11:32:55.324726Z","shell.execute_reply":"2024-06-11T11:32:55.460769Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"desc = pd.read_csv(BASE_DIR+'train_series_descriptions.csv')\ndesc.study_id = desc.study_id.astype(str)\ndesc.series_id = desc.series_id.astype(str)","metadata":{"execution":{"iopub.status.busy":"2024-06-11T11:32:55.464480Z","iopub.execute_input":"2024-06-11T11:32:55.465247Z","iopub.status.idle":"2024-06-11T11:32:55.497245Z","shell.execute_reply.started":"2024-06-11T11:32:55.465211Z","shell.execute_reply":"2024-06-11T11:32:55.496193Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"labels = pd.read_csv(BASE_DIR+'train.csv')\nlabels.study_id = labels.study_id.astype(str)","metadata":{"execution":{"iopub.status.busy":"2024-06-11T11:32:55.498770Z","iopub.execute_input":"2024-06-11T11:32:55.499212Z","iopub.status.idle":"2024-06-11T11:32:55.528230Z","shell.execute_reply.started":"2024-06-11T11:32:55.499172Z","shell.execute_reply":"2024-06-11T11:32:55.527146Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"conditions = np.unique(labels.columns[1:])\nclasses = []\nfor c in conditions:\n    classes.append(c+'_normal')\n    classes.append(c+'_moderate')\n    classes.append(c+'_severe')\nclasses_map = {classes[i]:i for i in range(len(classes))}\nclass_mapping = {i:classes[i] for i in range(len(classes))}\nN_CLASSES = len(class_mapping)","metadata":{"execution":{"iopub.status.busy":"2024-06-11T11:32:55.529620Z","iopub.execute_input":"2024-06-11T11:32:55.529992Z","iopub.status.idle":"2024-06-11T11:32:55.536586Z","shell.execute_reply.started":"2024-06-11T11:32:55.529963Z","shell.execute_reply":"2024-06-11T11:32:55.535502Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def prepare_labels(X):\n    out = np.zeros(N_CLASSES)\n    cols = X.index[1:]\n    X = X.values[1:]\n    \n    for x in cols[X=='Normal/Mild']: out[classes_map[x+'_normal']] = 1\n    for x in cols[X=='Moderate']: out[classes_map[x+'_moderate']] = 1\n    for x in cols[X=='Severe']: out[classes_map[x+'_severe']] = 1\n        \n    return out","metadata":{"execution":{"iopub.status.busy":"2024-06-11T11:32:55.538112Z","iopub.execute_input":"2024-06-11T11:32:55.538528Z","iopub.status.idle":"2024-06-11T11:32:55.550800Z","shell.execute_reply.started":"2024-06-11T11:32:55.538490Z","shell.execute_reply":"2024-06-11T11:32:55.549550Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def prepare_data(studies, df):\n    labels = []\n    image_paths = []\n    \n    for study_id in studies:\n        study_dir = TRAIN_DIR+study_id+'/'\n        sub_sample = df[df.study_id==study_id].fillna('Normal/Mild')\n        label = prepare_labels(sub_sample.iloc[0])\n        \n        for series_id in os.listdir(study_dir):\n            series_dir = study_dir+series_id+'/'\n            for z in os.listdir(series_dir):\n                z = z.split('.')[0]\n                #sub_desc = desc.where(desc.study_id==study_id).where(desc.series_id==series_id).dropna()\n                path = series_dir+z+'.dcm'\n                labels.append(label)\n                image_paths.append(path)\n    \n    data = tf.data.Dataset.from_tensor_slices((np.array(image_paths),np.array(labels, dtype='float32')))\n    return data","metadata":{"execution":{"iopub.status.busy":"2024-06-11T11:32:55.552143Z","iopub.execute_input":"2024-06-11T11:32:55.552559Z","iopub.status.idle":"2024-06-11T11:32:55.570045Z","shell.execute_reply.started":"2024-06-11T11:32:55.552521Z","shell.execute_reply":"2024-06-11T11:32:55.568886Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data = prepare_data(train_studies, labels)\nval_data = prepare_data(val_studies, labels)","metadata":{"execution":{"iopub.status.busy":"2024-06-11T11:32:55.571461Z","iopub.execute_input":"2024-06-11T11:32:55.572128Z","iopub.status.idle":"2024-06-11T11:33:44.727344Z","shell.execute_reply.started":"2024-06-11T11:32:55.572088Z","shell.execute_reply":"2024-06-11T11:33:44.726296Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_image(image_path):\n    raw_image = tf.io.read_file(image_path)\n    sp = tf.strings.split(tf.gather(tf.strings.split(image_path, 'images/'), 1), '/')\n    N = tf.size(sp)\n    LEN = tf.strings.length(tf.gather(sp, 0))+tf.strings.length(tf.gather(sp, 2))\n    \n    # Add missing file metadata to avoid warnnigs flooding\n    if   LEN==12: raw_image = tf.strings.regex_replace(raw_image, pattern=b'DICM\\x02\\x00\\x01\\x00', rewrite=b'DICM\\x02\\x00\\x00\\x00UL\\x04\\x00\\x92\\x00\\x00\\x00\\x02\\x00\\x01\\x00')\n    elif LEN==13: raw_image = tf.strings.regex_replace(raw_image, pattern=b'DICM\\x02\\x00\\x01\\x00', rewrite=b'DICM\\x02\\x00\\x00\\x00UL\\x04\\x00\\x92\\x00\\x00\\x00\\x02\\x00\\x01\\x00')\n    elif LEN==14: raw_image = tf.strings.regex_replace(raw_image, pattern=b'DICM\\x02\\x00\\x01\\x00', rewrite=b'DICM\\x02\\x00\\x00\\x00UL\\x04\\x00\\x94\\x00\\x00\\x00\\x02\\x00\\x01\\x00')\n    elif LEN==15: raw_image = tf.strings.regex_replace(raw_image, pattern=b'DICM\\x02\\x00\\x01\\x00', rewrite=b'DICM\\x02\\x00\\x00\\x00UL\\x04\\x00\\x94\\x00\\x00\\x00\\x02\\x00\\x01\\x00')\n    elif LEN==16: raw_image = tf.strings.regex_replace(raw_image, pattern=b'DICM\\x02\\x00\\x01\\x00', rewrite=b'DICM\\x02\\x00\\x00\\x00UL\\x04\\x00\\x96\\x00\\x00\\x00\\x02\\x00\\x01\\x00')\n    elif LEN==17: raw_image = tf.strings.regex_replace(raw_image, pattern=b'DICM\\x02\\x00\\x01\\x00', rewrite=b'DICM\\x02\\x00\\x00\\x00UL\\x04\\x00\\x96\\x00\\x00\\x00\\x02\\x00\\x01\\x00')\n    elif LEN==18: raw_image = tf.strings.regex_replace(raw_image, pattern=b'DICM\\x02\\x00\\x01\\x00', rewrite=b'DICM\\x02\\x00\\x00\\x00UL\\x04\\x00\\x98\\x00\\x00\\x00\\x02\\x00\\x01\\x00')\n    \n    img = tfio.image.decode_dicom_image(raw_image, scale='auto', dtype=tf.float32)\n    m, M=tf.math.reduce_min(img), tf.math.reduce_max(img)\n    img = (tf.image.grayscale_to_rgb(img)-m)/(M-m)\n    img = tf.image.resize(img, IMG_SIZE)[0]\n    return img\n\ndef load_dataset(image_path, labels):\n    image = load_image(image_path)\n    return {\"images\": tf.cast(image, tf.float32), \"labels\": tf.cast(labels, tf.float32)}","metadata":{"execution":{"iopub.status.busy":"2024-06-11T11:33:44.730412Z","iopub.execute_input":"2024-06-11T11:33:44.730768Z","iopub.status.idle":"2024-06-11T11:33:44.742386Z","shell.execute_reply.started":"2024-06-11T11:33:44.730737Z","shell.execute_reply":"2024-06-11T11:33:44.741104Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_ds = train_data.map(load_dataset, num_parallel_calls=tf.data.AUTOTUNE)\ntrain_ds = train_ds.ragged_batch(BATCH_SIZE, drop_remainder=True)\n\nval_ds = val_data.map(load_dataset, num_parallel_calls=tf.data.AUTOTUNE)\nval_ds = val_ds.ragged_batch(BATCH_SIZE, drop_remainder=True)","metadata":{"execution":{"iopub.status.busy":"2024-06-11T11:33:44.744047Z","iopub.execute_input":"2024-06-11T11:33:44.744394Z","iopub.status.idle":"2024-06-11T11:33:46.254416Z","shell.execute_reply.started":"2024-06-11T11:33:44.744364Z","shell.execute_reply":"2024-06-11T11:33:46.253314Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def dict_to_tuple(inputs):\n    return inputs[\"images\"], inputs[\"labels\"]\n\ntrain_ds = train_ds.map(dict_to_tuple, num_parallel_calls=tf.data.AUTOTUNE)\ntrain_ds = train_ds.prefetch(tf.data.AUTOTUNE)\n\nval_ds = val_ds.map(dict_to_tuple, num_parallel_calls=tf.data.AUTOTUNE)\nval_ds = val_ds.prefetch(tf.data.AUTOTUNE)","metadata":{"execution":{"iopub.status.busy":"2024-06-11T11:33:46.255994Z","iopub.execute_input":"2024-06-11T11:33:46.256414Z","iopub.status.idle":"2024-06-11T11:33:46.316100Z","shell.execute_reply.started":"2024-06-11T11:33:46.256377Z","shell.execute_reply":"2024-06-11T11:33:46.315107Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model ","metadata":{}},{"cell_type":"code","source":"backbone = keras_cv.models.EfficientNetV2Backbone.from_preset(PRETRAINED)\nmodel = keras.Sequential(\n    [\n        keras.layers.Input(shape=(None, None, 3)),\n        backbone,\n        keras.layers.GlobalMaxPooling2D(),\n        keras.layers.Dropout(rate=0.3),\n        keras.layers.Dense(N_CLASSES, activation=\"sigmoid\"),\n    ]\n)\nmodel.compile(optimizer=\"adam\",\n              loss=keras.losses.BinaryCrossentropy(),\n              metrics=[keras.metrics.AUC(name='auc')],\n             )\n\nmodel.summary()","metadata":{"execution":{"iopub.status.busy":"2024-06-11T11:33:46.317363Z","iopub.execute_input":"2024-06-11T11:33:46.317694Z","iopub.status.idle":"2024-06-11T11:34:13.580910Z","shell.execute_reply.started":"2024-06-11T11:33:46.317664Z","shell.execute_reply":"2024-06-11T11:34:13.579682Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import math\nimport matplotlib.pyplot as plt\n\n# Re-used the lr_scheduler from https://www.kaggle.com/code/awsaf49/birdclef24-kerascv-starter-train\ndef get_lr_callback(batch_size=8, mode='cos', epochs=10, plot=False):\n    lr_start, lr_max, lr_min = 5e-5, 8e-6 * batch_size, 1e-5\n    lr_ramp_ep, lr_sus_ep, lr_decay = 6, 0, 0.75\n\n    def lrfn(epoch):  # Learning rate update function\n        if epoch < lr_ramp_ep: lr = (lr_max - lr_start) / lr_ramp_ep * epoch + lr_start\n        elif epoch < lr_ramp_ep + lr_sus_ep: lr = lr_max\n        elif mode == 'exp': lr = (lr_max - lr_min) * lr_decay**(epoch - lr_ramp_ep - lr_sus_ep) + lr_min\n        elif mode == 'step': lr = lr_max * lr_decay**((epoch - lr_ramp_ep - lr_sus_ep) // 2)\n        elif mode == 'cos':\n            decay_total_epochs, decay_epoch_index = epochs - lr_ramp_ep - lr_sus_ep + 3, epoch - lr_ramp_ep - lr_sus_ep\n            phase = math.pi * decay_epoch_index / decay_total_epochs\n            lr = (lr_max - lr_min) * 0.5 * (1 + math.cos(phase)) + lr_min\n        return lr\n\n    if plot:  # Plot lr curve if plot is True\n        plt.figure(figsize=(10, 5))\n        plt.plot(np.arange(epochs), [lrfn(epoch) for epoch in np.arange(epochs)], marker='o')\n        plt.xlabel('epoch'); plt.ylabel('lr')\n        plt.title('LR Scheduler')\n        plt.show()\n\n    return keras.callbacks.LearningRateScheduler(lrfn, verbose=False)  # Create lr callback","metadata":{"execution":{"iopub.status.busy":"2024-06-11T11:34:13.582468Z","iopub.execute_input":"2024-06-11T11:34:13.583120Z","iopub.status.idle":"2024-06-11T11:34:13.594554Z","shell.execute_reply.started":"2024-06-11T11:34:13.583075Z","shell.execute_reply":"2024-06-11T11:34:13.593256Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"lr_cb = get_lr_callback(BATCH_SIZE, epochs=EPOCH, plot=True)","metadata":{"execution":{"iopub.status.busy":"2024-06-11T11:34:13.595955Z","iopub.execute_input":"2024-06-11T11:34:13.596309Z","iopub.status.idle":"2024-06-11T11:34:13.924379Z","shell.execute_reply.started":"2024-06-11T11:34:13.596279Z","shell.execute_reply":"2024-06-11T11:34:13.923253Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ckpt_cb = keras.callbacks.ModelCheckpoint(\"best_model.weights.h5\",\n                                         monitor='val_auc',\n                                         save_best_only=True,\n                                         save_weights_only=True,\n                                         mode='max')","metadata":{"execution":{"iopub.status.busy":"2024-06-11T11:34:13.926046Z","iopub.execute_input":"2024-06-11T11:34:13.926486Z","iopub.status.idle":"2024-06-11T11:34:13.933094Z","shell.execute_reply.started":"2024-06-11T11:34:13.926447Z","shell.execute_reply":"2024-06-11T11:34:13.931644Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"history = model.fit(\n    train_ds, \n    validation_data=val_ds, \n    epochs=EPOCH,\n    callbacks=[lr_cb, ckpt_cb], \n    verbose=1\n)","metadata":{"execution":{"iopub.status.busy":"2024-06-10T23:59:20.755330Z","iopub.execute_input":"2024-06-10T23:59:20.755734Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nauc = history.history['auc']\nval_auc = history.history['val_auc']\nloss = history.history['loss']\nval_loss = history.history['val_loss']\n\nepochs = range(len(loss))\nplt.plot(epochs, auc, 'r', label='Training auc')\nplt.plot(epochs, val_auc, 'b', label='Validation auc')\nplt.plot(epochs, loss, 'r', label='Training loss')\nplt.plot(epochs, val_loss, 'b', label='Validation loss')\nplt.legend(loc=0)\nplt.figure()\n\n\nplt.show()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}