{"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":"Hello Fellow Kagglers,\n\nThis notebook demonstrates the training process on generated data using EfficientNetV2-S as a backbone. 5 Models are trained with different seeds whose output is averaged during inference.\n\n14K generated signal samples and 7K noise samples are used for pretraining and the provided competition training dataset serves as a validation set. The model is fine tuned on this validation set for 1 epoch after pretraining on the generated data.\n\nThe model uses a single 360x256 spectogram as input, resized from 360x4096.\n\nAugmentations techniques involve cutout, horizontal/vertical flip and gaussion noise.\n\n[Generating Continuous Gravitational-Wave Sign PUB\n](https://www.kaggle.com/code/markwijkhuizen/generating-continuous-gravitational-wave-sign-pub)\n\n[Generating Continuous Gravitational-Wave Noise PUB\n](https://www.kaggle.com/code/markwijkhuizen/generating-continuous-gravitational-wave-noise-pub)\n\n[Inference](https://www.kaggle.com/code/markwijkhuizen/g2net-efficientnetv2-s-generated-data-tf-inference)\n\nV2: \n\n* H0(0.10, 0.01) -> H0(0.10, 0.04)\n* Reduced samples from 14K Signal/7K Noise -> 10K Signal/5K Noise\n* increased epochs 3 -> 5\n* ReLu -> LeakyReLu","metadata":{}},{"cell_type":"code","source":"import sys\nsys.path.append('/kaggle/input/efficientnetv2-pretrained-imagenet21k-weights/brain_automl/')\nsys.path.append('/kaggle/input/efficientnetv2-pretrained-imagenet21k-weights/brain_automl/efficientnetv2/')","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:22:52.616615Z","iopub.execute_input":"2022-11-22T17:22:52.617078Z","iopub.status.idle":"2022-11-22T17:22:52.648611Z","shell.execute_reply.started":"2022-11-22T17:22:52.616985Z","shell.execute_reply":"2022-11-22T17:22:52.647607Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport matplotlib as mpl\nimport tensorflow as tf\nimport tensorflow_addons as tfa\nimport effnetv2_model\n\nfrom tqdm.notebook import tqdm\nfrom sklearn.model_selection import train_test_split\n\nimport glob\nimport cv2\nimport gc\nimport os\nimport sys\nimport random\nimport math\nimport datetime\n\ntqdm.pandas()","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:22:52.653304Z","iopub.execute_input":"2022-11-22T17:22:52.6556Z","iopub.status.idle":"2022-11-22T17:22:58.119739Z","shell.execute_reply.started":"2022-11-22T17:22:52.655566Z","shell.execute_reply":"2022-11-22T17:22:58.118753Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Save Version","metadata":{}},{"cell_type":"code","source":"# Save datetime to easily determine notebook version from dataset\nnow = datetime.datetime.now().strftime(\"%d-%b-%Y %H-%M-%S\")\nnp.save(now, np.array([now]))","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:22:58.121068Z","iopub.execute_input":"2022-11-22T17:22:58.12174Z","iopub.status.idle":"2022-11-22T17:22:58.129032Z","shell.execute_reply.started":"2022-11-22T17:22:58.121701Z","shell.execute_reply":"2022-11-22T17:22:58.127976Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# MatplotLib Configuration","metadata":{}},{"cell_type":"code","source":"# MatplotLib Global Settings\nmpl.rcParams.update(mpl.rcParamsDefault)\nmpl.rcParams['xtick.labelsize'] = 16\nmpl.rcParams['ytick.labelsize'] = 16\nmpl.rcParams['axes.labelsize'] = 18\nmpl.rcParams['axes.titlesize'] = 24","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:22:58.131902Z","iopub.execute_input":"2022-11-22T17:22:58.1323Z","iopub.status.idle":"2022-11-22T17:22:58.139578Z","shell.execute_reply.started":"2022-11-22T17:22:58.132265Z","shell.execute_reply":"2022-11-22T17:22:58.13856Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Read Train DataFrame","metadata":{}},{"cell_type":"code","source":"# Read Train DataFrame\ntrain_labels = pd.read_csv('/kaggle/input/g2net-detecting-continuous-gravitational-waves/train_labels.csv')\n\ndisplay(train_labels.info())\n\ndisplay(train_labels.head())","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:22:58.142705Z","iopub.execute_input":"2022-11-22T17:22:58.143074Z","iopub.status.idle":"2022-11-22T17:22:58.185017Z","shell.execute_reply.started":"2022-11-22T17:22:58.143049Z","shell.execute_reply":"2022-11-22T17:22:58.184115Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Configuration","metadata":{}},{"cell_type":"code","source":"# Number of Samples in train dataset\nN_SAMPLES = len(train_labels)\n\n# Make 360x256 Patches\nTARGET_HEIGHT = 360\nTARGET_WIDTH = 256\nprint(f'TARGET_HEIGHT: {TARGET_HEIGHT}, TARGET_WIDTH: {TARGET_WIDTH}')\n\nBATCH_SIZE = 16\nBATCH_SIZE_FINE_TUNE = 8\nEPOCHS = 5\nEPOCHS_FINE_TUNE = 1\nNOISE_RATIO = 0.333\n\nLR_MAX = 4e-4\n\nIS_INTERACTIVE = os.environ['KAGGLE_KERNEL_RUN_TYPE'] == 'Interactive'\nVERBOSE = 1 if IS_INTERACTIVE else 2\n\nTRAIN_DIR = '/kaggle/input/g2net-detecting-continuous-gravitational-waves/train'\n\n# Generated Data Directories\nTRAIN_SAMPLES_DIR = '/kaggle/input/generating-continuous-gravitationalwave-sign-pub'\nTRAIN_SAMPLES_DIR_NOISE = '/kaggle/input/generating-continuous-gravitationalwave-noise-pub'\n\nINPUTS = ['H', 'L']\nN_INPUTS = len(INPUTS)","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:22:58.186095Z","iopub.execute_input":"2022-11-22T17:22:58.186364Z","iopub.status.idle":"2022-11-22T17:22:58.19489Z","shell.execute_reply.started":"2022-11-22T17:22:58.18634Z","shell.execute_reply":"2022-11-22T17:22:58.193464Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"# Determine number of samples using glob\nSAMPLE_IDXS = np.arange(len(glob.glob(f'{TRAIN_SAMPLES_DIR}/train_samples/x/*')))\nSAMPLE_IDXS_NOISE = np.arange(len(glob.glob(f'{TRAIN_SAMPLES_DIR_NOISE}/train_samples/x/*')))\nSAMPLE_IDXS_VAL = np.arange(len(glob.glob(f'{TRAIN_SAMPLES_DIR}/val_samples/x/*')))\n\n# Number of samples, each sample consist of a Hanford and Livington patch\nN_SAMPLE_IDXS = len(SAMPLE_IDXS)\nN_SAMPLE_IDXS_NOISE = len(SAMPLE_IDXS_NOISE)\nN_SAMPLE_IDXS_VAL = len(SAMPLE_IDXS_VAL)\n\n# Validation targets\nTARGETS_VAL = np.load(f'{TRAIN_SAMPLES_DIR}/TARGETS_VAL.npy')\n\n# Signal to noise ratios for generated signal samples\nSNRS = np.load(f'{TRAIN_SAMPLES_DIR}/SNRS.npy')\n\nprint(f'SAMPLE_IDXS shape: {SAMPLE_IDXS.shape}, SAMPLE_IDXS_VAL shape: {SAMPLE_IDXS_VAL.shape}')\nprint(f'SAMPLE_IDXS_NOISE shape: {SAMPLE_IDXS_NOISE.shape}')\nprint(f'TARGETS_VAL shape: {TARGETS_VAL.shape}')\nprint(f'SNRS shape: {SNRS.shape}')","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:22:58.197434Z","iopub.execute_input":"2022-11-22T17:22:58.198255Z","iopub.status.idle":"2022-11-22T17:22:58.684957Z","shell.execute_reply.started":"2022-11-22T17:22:58.198178Z","shell.execute_reply":"2022-11-22T17:22:58.683805Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Target Distribution Val\ndisplay(pd.Series(TARGETS_VAL).value_counts(normalize=True).to_frame('Ratio Val').round(3))","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:22:58.686622Z","iopub.execute_input":"2022-11-22T17:22:58.687195Z","iopub.status.idle":"2022-11-22T17:22:58.702796Z","shell.execute_reply.started":"2022-11-22T17:22:58.687156Z","shell.execute_reply":"2022-11-22T17:22:58.701649Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Sample Index to Number of Patches\nSAMPLE_IDX2N_PATCHES = dict([\n    (i, len(glob.glob(f'{TRAIN_SAMPLES_DIR}/train_samples/x/{i}/*.png'))) for i in tqdm(SAMPLE_IDXS)\n])\n\n# Sample Index to Number of Patches Noise\nSAMPLE_IDX2N_PATCHES_NOISE = dict([\n    (i, len(glob.glob(f'{TRAIN_SAMPLES_DIR_NOISE}/train_samples/x/{i}/*.png'))) for i in tqdm(SAMPLE_IDXS_NOISE)\n])\n\n# Validation Sample Index to Number of Patches\nSAMPLE_IDX2N_PATCHES_VAL = dict([\n    (i, len(glob.glob(f'{TRAIN_SAMPLES_DIR}/val_samples/x/{i}/*.png'))) for i in tqdm(SAMPLE_IDXS_VAL)\n])\n\n# Validate number of patches per sample, should be 2\ndisplay(pd.Series(SAMPLE_IDX2N_PATCHES.values()).value_counts().to_frame('Count Train Signal'))\ndisplay(pd.Series(SAMPLE_IDX2N_PATCHES_NOISE.values()).value_counts().to_frame('Count Train Noise'))\ndisplay(pd.Series(SAMPLE_IDX2N_PATCHES_VAL.values()).value_counts().to_frame('Count Val'))","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:22:58.704151Z","iopub.execute_input":"2022-11-22T17:22:58.704938Z","iopub.status.idle":"2022-11-22T17:23:39.594514Z","shell.execute_reply.started":"2022-11-22T17:22:58.704902Z","shell.execute_reply":"2022-11-22T17:23:39.593423Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Compute number of samples per type\nN_TRAIN_SAMPLES = 0\nN_TRAIN_SAMPLES_SIGNAL= 0\nN_VAL_SAMPLES = 0\n\nfor idx in SAMPLE_IDXS:\n    N_TRAIN_SAMPLES += SAMPLE_IDX2N_PATCHES[idx]\n    N_TRAIN_SAMPLES_SIGNAL += SAMPLE_IDX2N_PATCHES[idx]\n    \nfor idx in SAMPLE_IDXS_NOISE:\n    N_TRAIN_SAMPLES += SAMPLE_IDX2N_PATCHES_NOISE[idx]\n\nfor idx in SAMPLE_IDXS_VAL:\n    N_VAL_SAMPLES += SAMPLE_IDX2N_PATCHES_VAL[idx]\n\nprint(f'N_TRAIN_SAMPLES: {N_TRAIN_SAMPLES}, N_TRAIN_SAMPLES_SIGNAL: {N_TRAIN_SAMPLES_SIGNAL}')\nprint(f'N_VAL_SAMPLES: {N_VAL_SAMPLES}')","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:23:39.599157Z","iopub.execute_input":"2022-11-22T17:23:39.599753Z","iopub.status.idle":"2022-11-22T17:23:39.613633Z","shell.execute_reply.started":"2022-11-22T17:23:39.599725Z","shell.execute_reply":"2022-11-22T17:23:39.612362Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Samnple Index to Patch File Paths\nSAMPLE_IDX2FILE_PATHS = dict([\n    (i, glob.glob(f'{TRAIN_SAMPLES_DIR}/train_samples/x/{i}/*.png')) for i in tqdm(SAMPLE_IDXS)\n])\n\n# Samnple Index to Patch File Paths\nSAMPLE_IDX_NOISE2FILE_PATHS = dict([\n    (i, glob.glob(f'{TRAIN_SAMPLES_DIR_NOISE}/train_samples/x/{i}/*.png')) for i in tqdm(SAMPLE_IDXS_NOISE)\n])\n\n# Samnple Index to Patch File Paths VAL\nSAMPLE_IDX2FILE_PATHS_VAL = dict([\n    (i, glob.glob(f'{TRAIN_SAMPLES_DIR}/val_samples/x/{i}/*.png')) for i in tqdm(SAMPLE_IDXS_VAL)\n])","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:23:39.615281Z","iopub.execute_input":"2022-11-22T17:23:39.616163Z","iopub.status.idle":"2022-11-22T17:23:46.082435Z","shell.execute_reply.started":"2022-11-22T17:23:39.616123Z","shell.execute_reply":"2022-11-22T17:23:46.081436Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Training dataset will randomly select samples and apply augmentations\ndef get_train_dataset(bs, val=False, ret_snr=False):\n    X = np.zeros(shape=[bs, TARGET_HEIGHT, TARGET_WIDTH], dtype=np.uint8)\n    y = np.zeros(shape=[bs], dtype=np.int8)\n    snr = np.zeros(shape=[bs], dtype=np.float32)\n        \n    while True:\n        for b_idx in range(bs):\n            # Validation\n            if val:\n                index = int(np.random.choice(SAMPLE_IDXS_VAL, 1).squeeze())\n                file_path = random.choice(SAMPLE_IDX2FILE_PATHS_VAL[index])\n                y[b_idx] = TARGETS_VAL[index]\n            else:\n                # Generated Noise Sample\n                if random.random() < NOISE_RATIO:\n                    index = int(np.random.choice(SAMPLE_IDXS_NOISE, 1).squeeze())\n                    file_path = random.choice(SAMPLE_IDX_NOISE2FILE_PATHS[index])\n                    y[b_idx] = 0\n                    snr[b_idx] = 0\n                # Generated Signal Sample\n                else:\n                    index = int(np.random.choice(SAMPLE_IDXS, 1).squeeze())\n                    file_path = random.choice(SAMPLE_IDX2FILE_PATHS[index])\n                    y[b_idx] = 1\n                    snr[b_idx] = SNRS[index]\n                \n            \n                \n            # Load Image\n            X[b_idx] = cv2.imread(file_path, -1)\n\n            # Cutout\n            if np.random.rand() > 0.50:\n                X[b_idx] = np.flip(cv2.imread(file_path, -1), axis=1)\n                offset_x = np.random.randint(0, TARGET_HEIGHT // 2)\n                offset_y = np.random.randint(0, TARGET_WIDTH // 2)\n                size_x = TARGET_HEIGHT // 2\n                size_y = TARGET_WIDTH // 2\n                X[b_idx][offset_x:offset_x+size_x, offset_y:offset_y+size_y] = 0\n            \n            # Horizontal Flip\n            if np.random.rand() > 0.50:\n                X[b_idx] = np.flip(X[b_idx], axis=1)\n                \n            # Vertical Flip\n            if np.random.rand() > 0.50:\n                X[b_idx] = np.flip(X[b_idx], axis=0)\n        \n        if ret_snr:\n            yield X, y, snr\n        else:\n            yield X, y","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:23:46.084177Z","iopub.execute_input":"2022-11-22T17:23:46.084873Z","iopub.status.idle":"2022-11-22T17:23:46.098599Z","shell.execute_reply.started":"2022-11-22T17:23:46.084835Z","shell.execute_reply":"2022-11-22T17:23:46.097538Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Performance test and data statistics\ndef train_dataset_test():\n    train_dataset = get_train_dataset(BATCH_SIZE)\n    X, y = next(train_dataset)\n\n    print(f'X shape: {X.shape}, dtype: {X.dtype}', end=', ')\n    print(f'X mean: {X.mean():.2E}, std: {X.std():.2f}, min: {X.min():.2f}, max: {X.max():.2f}')\n\n    print(f'y: {y}, shape: {y.shape}, dtype: {y.dtype}')\n    \n    # Benchmark\n    N = 100\n    for index, _ in enumerate(tqdm(train_dataset, total=N)):\n        if index == N:\n            break\n            \n    return X, y\n    \nX_train_batch, y_train_batch = train_dataset_test()","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:23:46.10057Z","iopub.execute_input":"2022-11-22T17:23:46.101457Z","iopub.status.idle":"2022-11-22T17:23:56.99836Z","shell.execute_reply.started":"2022-11-22T17:23:46.10134Z","shell.execute_reply":"2022-11-22T17:23:56.99731Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Label Count\ndisplay(pd.Series(y_train_batch).value_counts().to_frame('Count'))","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:23:56.999763Z","iopub.execute_input":"2022-11-22T17:23:57.000357Z","iopub.status.idle":"2022-11-22T17:23:57.010587Z","shell.execute_reply.started":"2022-11-22T17:23:57.000318Z","shell.execute_reply":"2022-11-22T17:23:57.009621Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Augmentation test\nfor x in X_train_batch[:8]:\n    plt.imshow(x)\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:23:57.011902Z","iopub.execute_input":"2022-11-22T17:23:57.012844Z","iopub.status.idle":"2022-11-22T17:23:58.994152Z","shell.execute_reply.started":"2022-11-22T17:23:57.012809Z","shell.execute_reply":"2022-11-22T17:23:58.993261Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Validation set will iterate sequentially over all samples\ndef get_val_dataset(val=True, signal_only=False):\n    while True:\n        if val:\n            for index in SAMPLE_IDXS_VAL:\n                # Load target\n                y = np.expand_dims(TARGETS_VAL[index], axis=0)\n\n                for file_path in SAMPLE_IDX2FILE_PATHS_VAL[index]:\n                    X = cv2.imread(file_path, -1)\n                    X = np.expand_dims(X, axis=0)\n\n                    yield X, y\n        else:\n            for index in SAMPLE_IDXS_NOISE:\n                # Load target\n                y = np.array([[0]], dtype=np.int8)\n\n                for file_path in SAMPLE_IDX_NOISE2FILE_PATHS[index]:\n                    X = cv2.imread(file_path, -1)\n                    X = np.expand_dims(X, axis=0)\n\n                    yield X, y\n                    \n            for index in SAMPLE_IDXS:\n                # Load target\n                y = np.array([[1]], dtype=np.int8)\n\n                for file_path in SAMPLE_IDX2FILE_PATHS[index]:\n                    X = cv2.imread(file_path, -1)\n                    X = np.expand_dims(X, axis=0)\n\n                    yield X, y","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:23:58.995568Z","iopub.execute_input":"2022-11-22T17:23:59.00044Z","iopub.status.idle":"2022-11-22T17:23:59.012715Z","shell.execute_reply.started":"2022-11-22T17:23:59.000403Z","shell.execute_reply":"2022-11-22T17:23:59.011695Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Performance test and dataset statistics\ndef val_dataset_test():\n    val_dataset = get_val_dataset()\n    X, y = next(val_dataset)\n    \n    print(f'X shape: {X.shape}, dtype: {X.dtype}', end=', ')\n    print(f'X mean: {X.mean():.2E}, std: {X.std():.2f}, min: {X.min():.2f}, max: {X.max():.2f}')\n\n    print(f'y: {y}, shape: {y.shape}, dtype: {y.dtype}')\n    \n    # Benchmark\n    N = 100\n    for index, _ in enumerate(tqdm(val_dataset, total=N)):\n        if index == N:\n            return X, y\n    \nX_val_batch, y_val_batch = val_dataset_test()","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:23:59.01436Z","iopub.execute_input":"2022-11-22T17:23:59.015204Z","iopub.status.idle":"2022-11-22T17:23:59.705344Z","shell.execute_reply.started":"2022-11-22T17:23:59.01501Z","shell.execute_reply":"2022-11-22T17:23:59.704296Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.imshow(X_val_batch[0].squeeze())\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:23:59.706725Z","iopub.execute_input":"2022-11-22T17:23:59.70732Z","iopub.status.idle":"2022-11-22T17:23:59.950039Z","shell.execute_reply.started":"2022-11-22T17:23:59.707281Z","shell.execute_reply":"2022-11-22T17:23:59.949183Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train Validation Split","metadata":{}},{"cell_type":"code","source":"# Number of steps per epoch\nN_TRAIN_STEPS = N_TRAIN_SAMPLES // BATCH_SIZE + 1\nN_VAL_STEPS = N_VAL_SAMPLES\nprint(f'N_TRAIN_STEPS: {N_TRAIN_STEPS}, N_VAL_STEPS: {N_VAL_STEPS}')","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:23:59.951749Z","iopub.execute_input":"2022-11-22T17:23:59.952446Z","iopub.status.idle":"2022-11-22T17:23:59.958403Z","shell.execute_reply.started":"2022-11-22T17:23:59.952407Z","shell.execute_reply":"2022-11-22T17:23:59.957058Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"def get_model():\n    # Input spectogram 360x256\n    inputs = tf.keras.layers.Input(shape=[TARGET_HEIGHT, TARGET_WIDTH], dtype=tf.uint8)\n    \n    # Create 3-channel image\n    x = tf.expand_dims(inputs, axis=-1)\n    x = tf.tile(x, [1,1,1,3])\n    x = tf.cast(x, tf.float32)\n    # Imagenet normalization\n    x = tf.keras.applications.imagenet_utils.preprocess_input(x, mode='torch')\n    # Gaussion noise\n    x = tf.keras.layers.GaussianNoise(0.10)(x)\n\n    # EfficientNetV2-S backbone\n    cnn = effnetv2_model.get_model(f'efficientnetv2-s', include_top=False, weights='imagenet21k-ft1k')\n\n    x = cnn(x)\n    # Apply ReLu activation, there should not be a negative (noise) filter\n    # Noise filters do not make sense...\n    x = tf.keras.layers.LeakyReLU()(x)\n    # Dropout for regularization\n    x = tf.keras.layers.Dropout(0.30)(x)\n\n    # Output is a single neuron with sigmoid activation initialized with HE weights\n    output = tf.keras.layers.Dense(1, activation='sigmoid', kernel_initializer='he_uniform')(x)\n    \n    # Model\n    model = tf.keras.models.Model(inputs=inputs, outputs=output)\n    \n    # LOSS\n    loss = tf.keras.losses.BinaryCrossentropy(from_logits=False)\n\n    # OPTIMIZER\n    optimizer = tfa.optimizers.AdamW(learning_rate=LR_MAX, weight_decay=LR_MAX*1e-2, epsilon=1e-7)\n\n    # METRICS\n    metrics = [\n        tf.keras.metrics.BinaryAccuracy(),\n        tf.keras.metrics.Precision(),\n        tf.keras.metrics.Recall(),\n        tf.keras.metrics.AUC(),\n    ]\n\n    # Compile Model\n    model.compile(optimizer=optimizer, loss=loss, metrics=metrics)\n    \n    return model","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:23:59.960173Z","iopub.execute_input":"2022-11-22T17:23:59.960899Z","iopub.status.idle":"2022-11-22T17:23:59.972413Z","shell.execute_reply.started":"2022-11-22T17:23:59.960861Z","shell.execute_reply":"2022-11-22T17:23:59.971396Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tf.keras.backend.clear_session()\ngc.collect()\n\nmodel = get_model()","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:23:59.97538Z","iopub.execute_input":"2022-11-22T17:23:59.975853Z","iopub.status.idle":"2022-11-22T17:24:38.036079Z","shell.execute_reply.started":"2022-11-22T17:23:59.975819Z","shell.execute_reply":"2022-11-22T17:24:38.035092Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Model Architecture Summary\nprint(model.summary())","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:24:38.03765Z","iopub.execute_input":"2022-11-22T17:24:38.038014Z","iopub.status.idle":"2022-11-22T17:24:38.073507Z","shell.execute_reply.started":"2022-11-22T17:24:38.037976Z","shell.execute_reply":"2022-11-22T17:24:38.072409Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Model Architecture Visualization\ntf.keras.utils.plot_model(model, show_shapes=True, show_dtype=True, show_layer_names=True, expand_nested=False)","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:24:38.07484Z","iopub.execute_input":"2022-11-22T17:24:38.075302Z","iopub.status.idle":"2022-11-22T17:24:38.984361Z","shell.execute_reply.started":"2022-11-22T17:24:38.075266Z","shell.execute_reply":"2022-11-22T17:24:38.983209Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Weight Initialization Test","metadata":{}},{"cell_type":"code","source":"# Validation Baseline, sanity check, should be random guessing\n_ = model.evaluate(\n        get_val_dataset(val=True),\n        steps=N_VAL_STEPS,\n        verbose=VERBOSE,\n    )","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:24:38.986447Z","iopub.execute_input":"2022-11-22T17:24:38.987137Z","iopub.status.idle":"2022-11-22T17:25:28.687563Z","shell.execute_reply.started":"2022-11-22T17:24:38.987094Z","shell.execute_reply":"2022-11-22T17:25:28.686508Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Training Output Baseline, verify output is ~0.50 centered for correct gradient flow\ntrain_preds = model.predict(\n        get_val_dataset(val=False),\n        steps=1024,\n        verbose=VERBOSE,\n    ).squeeze()","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:25:28.689506Z","iopub.execute_input":"2022-11-22T17:25:28.689936Z","iopub.status.idle":"2022-11-22T17:25:59.243384Z","shell.execute_reply.started":"2022-11-22T17:25:28.689895Z","shell.execute_reply":"2022-11-22T17:25:59.242408Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Visualize Baseline Predictions\nplt.figure(figsize=(15,8))\nplt.title(f'Train Prediction Initialized Model')\npd.Series(train_preds).plot(kind='hist')\nplt.xticks(np.arange(0, 1.1, 0.1))\nplt.grid()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:25:59.24501Z","iopub.execute_input":"2022-11-22T17:25:59.245326Z","iopub.status.idle":"2022-11-22T17:25:59.634896Z","shell.execute_reply.started":"2022-11-22T17:25:59.245299Z","shell.execute_reply":"2022-11-22T17:25:59.633865Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Learning Rate Scheduler","metadata":{}},{"cell_type":"code","source":"# Cosine Decay with exponential warmup\ndef lrfn(current_step, num_warmup_steps, lr_max, num_cycles=0.50, num_training_steps=EPOCHS):\n    \n    if current_step < num_warmup_steps:\n        return lr_max * 0.50 ** (num_warmup_steps - current_step)\n    else:\n        progress = float(current_step - num_warmup_steps) / float(max(1, num_training_steps - num_warmup_steps))\n\n        return max(0.0, 0.5 * (1.0 + math.cos(math.pi * float(num_cycles) * 2.0 * progress))) * lr_max","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:25:59.639846Z","iopub.execute_input":"2022-11-22T17:25:59.642266Z","iopub.status.idle":"2022-11-22T17:25:59.651394Z","shell.execute_reply.started":"2022-11-22T17:25:59.642211Z","shell.execute_reply":"2022-11-22T17:25:59.65045Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Plots the learning rate scheduler\ndef plot_lr_schedule(lr_schedule, epochs):\n    fig = plt.figure(figsize=(20, 10))\n    plt.plot([None] + lr_schedule + [None])\n    # X Labels\n    x = np.arange(1, epochs + 1)\n    x_axis_labels = [i if epochs <= 40 or i % 5 == 0 or i == 1 else None for i in range(1, epochs + 1)]\n    plt.xlim([1, epochs])\n    plt.xticks(x, x_axis_labels) # set tick step to 1 and let x axis start at 1\n    \n    # Increase y-limit for better readability\n    plt.ylim([0, max(lr_schedule) * 1.1])\n    \n    # Title\n    schedule_info = f'start: {lr_schedule[0]:.1E}, max: {max(lr_schedule):.1E}, final: {lr_schedule[-1]:.1E}'\n    plt.title(f'Step Learning Rate Schedule, {schedule_info}', size=18, pad=12)\n    \n    # Plot Learning Rates\n    for x, val in enumerate(lr_schedule):\n        if epochs <= 40 or x % 5 == 0 or x is epochs - 1:\n            if x < len(lr_schedule) - 1:\n                if lr_schedule[x - 1] < val:\n                    ha = 'right'\n                else:\n                    ha = 'left'\n            elif x == 0:\n                ha = 'right'\n            else:\n                ha = 'left'\n            plt.plot(x + 1, val, 'o', color='black');\n            offset_y = (max(lr_schedule) - min(lr_schedule)) * 0.02\n            plt.annotate(f'{val:.1E}', xy=(x + 1, val + offset_y), size=12, ha=ha)\n    \n    plt.xlabel('Epoch', size=16, labelpad=5)\n    plt.ylabel('Learning Rate', size=16, labelpad=5)\n    plt.grid()\n    plt.show()\n\n# Learning rate for encoder\nLR_SCHEDULE = [lrfn(step, num_warmup_steps=0, lr_max=LR_MAX, num_cycles=0.50) for step in range(EPOCHS)]\nplot_lr_schedule(LR_SCHEDULE, epochs=EPOCHS)","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:25:59.663188Z","iopub.execute_input":"2022-11-22T17:25:59.665444Z","iopub.status.idle":"2022-11-22T17:26:00.079138Z","shell.execute_reply.started":"2022-11-22T17:25:59.665386Z","shell.execute_reply":"2022-11-22T17:26:00.078161Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Learning Rate Callback\nlr_callback = tf.keras.callbacks.LearningRateScheduler(lambda step: LR_SCHEDULE[step], verbose=1)","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:26:00.080492Z","iopub.execute_input":"2022-11-22T17:26:00.08149Z","iopub.status.idle":"2022-11-22T17:26:00.0871Z","shell.execute_reply.started":"2022-11-22T17:26:00.081447Z","shell.execute_reply":"2022-11-22T17:26:00.085716Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Weight Decay Callback","metadata":{}},{"cell_type":"code","source":"class WeightDecayCallback(tf.keras.callbacks.Callback):\n    def __init__(self, wd_ratio=0.01):\n        self.step_counter = 0\n        self.wd_ratio = wd_ratio\n    \n    def on_epoch_begin(self, epoch, logs=None):\n        model.optimizer.weight_decay = model.optimizer.learning_rate * self.wd_ratio\n        print(f'learning rate: {model.optimizer.learning_rate.numpy():.2e}, weight decay: {model.optimizer.weight_decay.numpy():.2e}')","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:32:46.845863Z","iopub.execute_input":"2022-11-22T17:32:46.846268Z","iopub.status.idle":"2022-11-22T17:32:46.855137Z","shell.execute_reply.started":"2022-11-22T17:32:46.846217Z","shell.execute_reply":"2022-11-22T17:32:46.854201Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"code","source":"# Function to seed everything for deterministic results\ndef set_seeds(seed):\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    random.seed(seed)\n    tf.random.set_seed(seed)\n    np.random.seed(seed)","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:32:47.558682Z","iopub.execute_input":"2022-11-22T17:32:47.559085Z","iopub.status.idle":"2022-11-22T17:32:47.565393Z","shell.execute_reply.started":"2022-11-22T17:32:47.559053Z","shell.execute_reply":"2022-11-22T17:32:47.563943Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"N_FOLDS = 5\nfor fold in range(N_FOLDS):\n    print('='*10, f' FOLD {fold} ', '='*10)\n    \n    # Set seed differently each training run\n    set_seeds(fold)\n    \n    # Clean Up all previour models\n    tf.keras.backend.clear_session()\n    gc.collect()\n\n    model = get_model()\n\n    # Checkpoint callback\n    model_checkpoint_callback = tf.keras.callbacks.ModelCheckpoint(\n        f'fold_{fold}_model' + '_epochs-{epoch:02d}_val_auc-{val_auc:.3f}.hdf5',\n        monitor = 'val_auc',\n        verbose = 1,\n        save_best_only = False,\n        save_weights_only = True,\n    )\n\n    # Train Model\n    history = model.fit(\n            get_train_dataset(BATCH_SIZE),\n            steps_per_epoch=N_TRAIN_SAMPLES // BATCH_SIZE,\n            validation_data=get_val_dataset(),\n            validation_steps=N_VAL_STEPS,\n            epochs = EPOCHS,\n            verbose = VERBOSE,\n            callbacks = [\n                lr_callback,\n                model_checkpoint_callback,\n                WeightDecayCallback(0.01),\n            ],\n        )\n    \n    \n    model.save_weights(f'fold_{fold}_model_pretrained.h5')\n\n    # Fine Tune Model on provided competition training set\n    model.fit(\n            get_train_dataset(BATCH_SIZE_FINE_TUNE, val=True),\n            steps_per_epoch=N_VAL_SAMPLES // BATCH_SIZE_FINE_TUNE,\n            epochs = EPOCHS_FINE_TUNE,\n            verbose = VERBOSE,\n            callbacks = [],\n        )\n\n    model.save_weights(f'fold_{fold}_model_fine_tuned.h5')","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:32:48.153604Z","iopub.execute_input":"2022-11-22T17:32:48.153985Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Validation Predictions Last Model","metadata":{}},{"cell_type":"code","source":"# Validation predictions of last\nval_y_true = []\n\nfor i, (_, target) in enumerate(get_val_dataset()):\n    if i == N_VAL_STEPS:\n        break\n        \n    val_y_true.append(target)\n    \nval_y_true = np.array(val_y_true, dtype=np.int8).squeeze()\nprint(f'val_y_true shape: {val_y_true.shape}')","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:32:44.677783Z","iopub.status.idle":"2022-11-22T17:32:44.679654Z","shell.execute_reply.started":"2022-11-22T17:32:44.67937Z","shell.execute_reply":"2022-11-22T17:32:44.679396Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"display(pd.Series(val_y_true).value_counts().to_frame('Count'))","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:32:44.681264Z","iopub.status.idle":"2022-11-22T17:32:44.682068Z","shell.execute_reply.started":"2022-11-22T17:32:44.681805Z","shell.execute_reply":"2022-11-22T17:32:44.681829Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Validation Prediction Pretrained Model\nmodel.load_weights(f'fold_{N_FOLDS-1}_model_pretrained.h5')\nval_preds = model.predict(\n        get_val_dataset(),\n        steps=N_VAL_STEPS,\n        verbose=VERBOSE,\n    ).squeeze()","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:32:44.683603Z","iopub.status.idle":"2022-11-22T17:32:44.684394Z","shell.execute_reply.started":"2022-11-22T17:32:44.684117Z","shell.execute_reply":"2022-11-22T17:32:44.684141Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"display(pd.Series(val_preds).describe().to_frame('Value'))","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:32:44.685848Z","iopub.status.idle":"2022-11-22T17:32:44.686641Z","shell.execute_reply.started":"2022-11-22T17:32:44.686369Z","shell.execute_reply":"2022-11-22T17:32:44.686393Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(15,8))\nplt.title(f'Val Prediction Trained Model')\npd.Series(val_preds).plot(kind='hist')\nplt.xlim(0, 1)\nplt.xticks(np.arange(0, 1.1, 0.1))\nplt.grid()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:32:44.688074Z","iopub.status.idle":"2022-11-22T17:32:44.6889Z","shell.execute_reply.started":"2022-11-22T17:32:44.688644Z","shell.execute_reply":"2022-11-22T17:32:44.688668Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(15,8))\nplt.title(f'Val Prediction Error Trained Model')\npd.Series(val_y_true - val_preds).plot(kind='hist', bins=100)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:32:44.690328Z","iopub.status.idle":"2022-11-22T17:32:44.691103Z","shell.execute_reply.started":"2022-11-22T17:32:44.690848Z","shell.execute_reply":"2022-11-22T17:32:44.690872Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train Predictions","metadata":{}},{"cell_type":"code","source":"train_y_true = []\n\nfor i, (_, target) in tqdm(enumerate(get_val_dataset(val=False)), total=N_TRAIN_SAMPLES):\n    if i == N_TRAIN_SAMPLES:\n        break\n        \n    train_y_true.append(target)\n    \ntrain_y_true = np.array(train_y_true, dtype=np.int8).squeeze()\nprint(f'train_y_true shape: {train_y_true.shape}')","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:32:44.692536Z","iopub.status.idle":"2022-11-22T17:32:44.693308Z","shell.execute_reply.started":"2022-11-22T17:32:44.693038Z","shell.execute_reply":"2022-11-22T17:32:44.693062Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"display(pd.Series(train_y_true).value_counts().to_frame('Count'))","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:32:44.694712Z","iopub.status.idle":"2022-11-22T17:32:44.695506Z","shell.execute_reply.started":"2022-11-22T17:32:44.695228Z","shell.execute_reply":"2022-11-22T17:32:44.695269Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Training Output Pretrained Model\ntrain_preds = model.predict(\n        get_val_dataset(val=False),\n        steps=1024,\n        verbose=VERBOSE,\n    ).squeeze()","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:32:44.69691Z","iopub.status.idle":"2022-11-22T17:32:44.697705Z","shell.execute_reply.started":"2022-11-22T17:32:44.697437Z","shell.execute_reply":"2022-11-22T17:32:44.697461Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(15,8))\nplt.title(f'Train Prediction Trained Model')\npd.Series(train_preds).plot(kind='hist')\nplt.xlim(0, 1)\nplt.xticks(np.arange(0, 1.1, 0.1))\nplt.grid()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:32:44.699097Z","iopub.status.idle":"2022-11-22T17:32:44.699891Z","shell.execute_reply.started":"2022-11-22T17:32:44.699626Z","shell.execute_reply":"2022-11-22T17:32:44.69965Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# AUC By SNR\n\nProvides the AUC by signal-to-noise ratio of the pretrained model to get a sense of how well the model can fit each signal-to-noise ratio.","metadata":{}},{"cell_type":"code","source":"PERCENTILES = np.arange(0.00, 1.00, 0.1)","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:32:44.701292Z","iopub.status.idle":"2022-11-22T17:32:44.702073Z","shell.execute_reply.started":"2022-11-22T17:32:44.701813Z","shell.execute_reply":"2022-11-22T17:32:44.701837Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def show_snr_percentiles():\n    # Load Signal to Noise Ratios and Filter Noise Samples\n    SNRS = pd.Series(np.load(f'{TRAIN_SAMPLES_DIR}/SNRS.npy'))\n    # Show Percentiles\n    display(SNRS.describe(percentiles=PERCENTILES).to_frame('Value'))\n    \nshow_snr_percentiles()","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:32:44.703503Z","iopub.status.idle":"2022-11-22T17:32:44.704288Z","shell.execute_reply.started":"2022-11-22T17:32:44.704012Z","shell.execute_reply":"2022-11-22T17:32:44.704037Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"y_true = []\ny_pred = []\nsnr = []\nN = 32\n\n# Train Predictions Signal\nfor idx, (x, y, s) in tqdm(enumerate(get_train_dataset(128, ret_snr=True)), total=N):\n    if idx == N:\n        break\n    else:\n        y_true += y.tolist()\n        y_pred += model.predict_on_batch(x).squeeze().tolist()\n        snr += s.tolist()\n        \ny_true = np.array(y_true, dtype=np.int8)\ny_pred = np.array(y_pred, dtype=np.float32)\nsnr = np.array(snr, dtype=np.float32)\n        \nprint(f'y_true shape: {y_true.shape}, y_pred shape: {y_pred.shape}, snr shape: {snr.shape}')\ndisplay(pd.Series(y_true).value_counts().to_frame('count'))","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:32:44.705675Z","iopub.status.idle":"2022-11-22T17:32:44.706453Z","shell.execute_reply.started":"2022-11-22T17:32:44.706179Z","shell.execute_reply":"2022-11-22T17:32:44.706202Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_pred_df = pd.DataFrame({\n    'y_true': y_true,\n    'y_pred': y_pred,\n    'snr': snr,\n})\n\n# Split Validation Sets\nSNR_PERCENTILES = np.percentile(SNRS, np.concatenate((PERCENTILES * 100, (PERCENTILES * 100) + 10))).reshape([2, -1]).T\nSNR_BUCKETS = []\nfor snr_idx, (snr_min, snr_max) in enumerate(SNR_PERCENTILES):\n    print(f'bucket {snr_idx} | snr min: {snr_min:.0f}, max: {snr_max:.0f}')\n    SNR_BUCKETS.append((\n        f'{10*snr_idx}p_{10*(snr_idx + 1)}p snr_{int(snr_min)}_{int(snr_max)}',\n        snr_min,\n        snr_max,\n    ))\n\ndef get_snr_bucket(snr):\n    # put noise in random bucker\n    if snr == 0:\n        return np.random.choice([t[0] for t in SNR_BUCKETS], 1).squeeze()\n    \n    for name, snr_min, snr_max in SNR_BUCKETS:\n        if snr >= snr_min and snr < snr_max:\n            return name\n\ntrain_pred_df['snr_bucket'] = train_pred_df['snr'].apply(get_snr_bucket)\n\ndisplay(train_pred_df.head())","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:32:44.707859Z","iopub.status.idle":"2022-11-22T17:32:44.708654Z","shell.execute_reply.started":"2022-11-22T17:32:44.708384Z","shell.execute_reply":"2022-11-22T17:32:44.708408Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"display(train_pred_df.groupby('snr_bucket')['y_true'].value_counts().to_frame())","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:32:44.710043Z","iopub.status.idle":"2022-11-22T17:32:44.71084Z","shell.execute_reply.started":"2022-11-22T17:32:44.710589Z","shell.execute_reply":"2022-11-22T17:32:44.710614Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Verify even distribution of positive and negative samples per snr bucket\ndisplay(\n   (\n       train_pred_df.groupby('snr_bucket')['y_true'].value_counts(normalize=True).to_frame()\n        * 100\n   ).round(1)\n)","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:32:44.712223Z","iopub.status.idle":"2022-11-22T17:32:44.71302Z","shell.execute_reply.started":"2022-11-22T17:32:44.712757Z","shell.execute_reply":"2022-11-22T17:32:44.712781Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def auc_by_snr(g):\n    return tf.keras.metrics.AUC()(g['y_true'], g['y_pred']).numpy()\n\n# AUC by signal-to-noise ratio bucket\ndisplay(train_pred_df.groupby('snr_bucket').apply(auc_by_snr).to_frame('AUC').sort_index().round(3))","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:32:44.714411Z","iopub.status.idle":"2022-11-22T17:32:44.715183Z","shell.execute_reply.started":"2022-11-22T17:32:44.714925Z","shell.execute_reply":"2022-11-22T17:32:44.714949Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Signal to noise ratios above ~75 can be perfectly fitted\nauc_by_snr_df =  train_pred_df.groupby('snr_bucket').apply(auc_by_snr).to_frame('AUC').sort_index()\n\nplt.figure(figsize=(15,10))\nplt.title('AUC By SNR Bucket')\nplt.plot(auc_by_snr_df['AUC'])\nplt.yticks(np.arange(0.0, 1.1, 0.1))\nplt.xticks(rotation=45)\nplt.ylabel('AUC')\nplt.ylim(0,1)\nplt.grid()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:32:44.716639Z","iopub.status.idle":"2022-11-22T17:32:44.717427Z","shell.execute_reply.started":"2022-11-22T17:32:44.717153Z","shell.execute_reply":"2022-11-22T17:32:44.717177Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training History","metadata":{}},{"cell_type":"code","source":"# Plots training history metrics\ndef plot_history_metric(metric, f_best=np.argmax, ylim=None, yscale=None, yticks=None):\n    plt.figure(figsize=(20, 10))\n    \n    values = history.history[metric]\n    N_EPOCHS = len(values)\n    val = 'val' in ''.join(history.history.keys())\n    # Epoch Ticks\n    if N_EPOCHS <= 20:\n        x = np.arange(1, N_EPOCHS + 1)\n    else:\n        x = [1, 5] + [10 + 5 * idx for idx in range((N_EPOCHS - 10) // 5 + 1)]\n\n    x_ticks = np.arange(1, N_EPOCHS+1)\n\n    # Validation\n    if val:\n        val_values = history.history[f'val_{metric}']\n        val_argmin = f_best(val_values)\n        plt.plot(x_ticks, val_values, label=f'val')\n\n    # summarize history for accuracy\n    plt.plot(x_ticks, values, label=f'train')\n    argmin = f_best(values)\n    plt.scatter(argmin + 1, values[argmin], color='red', s=75, marker='o', label=f'train_best')\n    if val:\n        plt.scatter(val_argmin + 1, val_values[val_argmin], color='purple', s=75, marker='o', label=f'val_best')\n\n    plt.title(f'Model {metric}', fontsize=24, pad=10)\n    plt.ylabel(metric, fontsize=20, labelpad=10)\n\n    if ylim:\n        plt.ylim(ylim)\n\n    if yscale is not None:\n        plt.yscale(yscale)\n        \n    if yticks is not None:\n        plt.yticks(yticks, fontsize=16)\n\n    plt.xlabel('epoch', fontsize=20, labelpad=10)        \n    plt.tick_params(axis='x', labelsize=8)\n    plt.xticks(x, fontsize=16) # set tick step to 1 and let x axis start at 1\n    plt.yticks(fontsize=16)\n    \n    plt.legend(prop={'size': 10})\n    plt.grid()\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:32:44.718875Z","iopub.status.idle":"2022-11-22T17:32:44.719688Z","shell.execute_reply.started":"2022-11-22T17:32:44.719421Z","shell.execute_reply":"2022-11-22T17:32:44.719445Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_history_metric('loss', f_best=np.argmin)","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:32:44.721092Z","iopub.status.idle":"2022-11-22T17:32:44.721878Z","shell.execute_reply.started":"2022-11-22T17:32:44.72162Z","shell.execute_reply":"2022-11-22T17:32:44.721643Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_history_metric('binary_accuracy', ylim=[0,1])","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:32:44.723348Z","iopub.status.idle":"2022-11-22T17:32:44.72416Z","shell.execute_reply.started":"2022-11-22T17:32:44.723899Z","shell.execute_reply":"2022-11-22T17:32:44.723924Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_history_metric('precision', ylim=[0,1])","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:32:44.725611Z","iopub.status.idle":"2022-11-22T17:32:44.726425Z","shell.execute_reply.started":"2022-11-22T17:32:44.726139Z","shell.execute_reply":"2022-11-22T17:32:44.726162Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_history_metric('recall', ylim=[0,1])","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:32:44.727853Z","iopub.status.idle":"2022-11-22T17:32:44.728687Z","shell.execute_reply.started":"2022-11-22T17:32:44.728425Z","shell.execute_reply":"2022-11-22T17:32:44.72845Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_history_metric('auc', ylim=[0,1])","metadata":{"execution":{"iopub.status.busy":"2022-11-22T17:32:44.730109Z","iopub.status.idle":"2022-11-22T17:32:44.730935Z","shell.execute_reply.started":"2022-11-22T17:32:44.730639Z","shell.execute_reply":"2022-11-22T17:32:44.730662Z"},"trusted":true},"execution_count":null,"outputs":[]}]}