{"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":"# G2NET: Pytorch+Generated Realistic Noise\n\n\n## TL;DR\n\nAs it was found in [**G2Net: Winning Strategy with External Data**\n](https://www.kaggle.com/code/vslaykovsky/g2net-winning-strategy-with-external-data), 20% of all test samples were generated from data comming from real detectors L1 and H1. E.g. organizers took real samples of real noise and overlayed generated signals. Unfortunately, the training set only contains generated noise, so there is an inevitable distribution shift between train and test samples. This notebook targets this specific problem.\n\nIn this notebook we use noise generated in [**G2NET: Realistic Simulation of Test Noise**](https://www.kaggle.com/code/vslaykovsky/g2net-realistic-simulation-of-test-noise), combine it with [pure signal](https://www.kaggle.com/code/vslaykovsky/g2net-generating-pure-signal) to produce synthetic data.\nThis synthetic data is then used to train a better model for the 20% of \"real\" test samples. \n\nThe output is then combined with the best public solution to improve LB score. \n\n\n## Notes:\n\n* For details on how noise and signal are combined refer to `RealisticNoiseDataset`. We use a simple weighted sum. Additionally we use gaussian noise to mitigate overfitting. \n* Optuna is used for hyperparameter tuning. \n* 5 folds model is trained for ensemble inference. \n\n\n**Smash that like button and subscribe for more eye-popping notebooks!**","metadata":{}},{"cell_type":"markdown","source":"## Utils","metadata":{}},{"cell_type":"code","source":"from kaggle_secrets import UserSecretsClient\nuser_secrets = UserSecretsClient()\nsecret_value_0 = user_secrets.get_secret(\"WANDB_API_KEY\")","metadata":{"execution":{"iopub.status.busy":"2022-12-12T14:35:13.548109Z","iopub.execute_input":"2022-12-12T14:35:13.548813Z","iopub.status.idle":"2022-12-12T14:35:13.641404Z","shell.execute_reply.started":"2022-12-12T14:35:13.548777Z","shell.execute_reply":"2022-12-12T14:35:13.640436Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#!pip install -q timm\nimport os\ntry:\n    from kaggle_secrets import UserSecretsClient\n    user_secrets = UserSecretsClient()\n    WANDB_API_KEY = user_secrets.get_secret(\"WANDB_API_KEY\")\n    os.environ['WANDB_API_KEY'] = WANDB_API_KEY\n#     NOISE_DIR = '/kaggle/input/g2net-realistic-simulation-of-test-noise/data/realistic_noise/images/'\n    NOISE_DIR = '/kaggle/input/realistic-noise-256/data/realistic_noise/images/'\n    \n#     SIGNAL_DIR = '/kaggle/input/g2net-generating-pure-signal/data/pure_signal'\n    SIGNAL_DIR = '/kaggle/input/g2net-pure-signal/pure_signal'    \n    !pip install -q timm\n    !pip install -q git+https://github.com/PyFstat/PyFstat@python37\nexcept:\n    print('Running locally')\n    NOISE_DIR = 'data/realistic_noise/images/'\n    SIGNAL_DIR = 'data/pure_signal/'\n    os.environ['WANDB_API_KEY'] = 'your_key_here'\n\nfrom timm.data.transforms_factory import create_transform\nimport torchvision\nfrom torch.utils.data import Dataset\nimport glob\nimport numpy as np\nfrom PIL import Image\nimport matplotlib.pyplot as plt\nimport timm\nimport torch\nfrom sklearn.metrics import *\nfrom tqdm.notebook import tqdm\nimport gc\nimport wandb\nimport os\nimport pandas as pd\nimport re\n\nBATCH_SIZE = 32\nDEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\n    \nPOSITIVE_RATE = 0.5\nSIGNAL_LOW = 0.02\nSIGNAL_HIGH = 0.10\nEVAL_PCT = 0.2\nDEBUG = True\nOPTUNA = False\nTRAIN = True\nFOLDS = [0, 1, 2, 3, 4]\n# FOLDS = [0] \nN_FOLDS = 5\n\n# class Config:\n#finished with value: 0.9110377703633762 \n#parameters: \n#{'LR': 0.00023587771289046934, 'DROPOUT': 0.2, 'MAX_GRAD_NORM': 16.068213932163527, \n#'EPOCHS': 7.0, 'GAUSSIAN_NOISE': 1.0, 'ONE_CYCLE_PCT_START': 0.0, \n#'MODEL': 'efficientnetv2_rw_s', 'ONE_CYCLE': True}. \nLR = 0.00056\nDROPOUT = 0.25#0.25\nMAX_GRAD_NORM = 1.36\nEPOCHS = 3\nGAUSSIAN_NOISE = 2.#2.\nONE_CYCLE_PCT_START=0.1\nMODEL = 'inception_v4'#'efficientnetv2_rw_s'#\nONE_CYCLE = True\n\nWANDB_RUN = f'{MODEL}-LR:{LR}-DR:{DROPOUT}-E:{EPOCHS}-MGN:{MAX_GRAD_NORM}-NOISE:{GAUSSIAN_NOISE}-OC:{int(ONE_CYCLE)}-OCPS:{ONE_CYCLE_PCT_START}'\nWANDB_RUN","metadata":{"execution":{"iopub.status.busy":"2022-12-12T14:35:13.643566Z","iopub.execute_input":"2022-12-12T14:35:13.643959Z","iopub.status.idle":"2022-12-12T14:35:34.658479Z","shell.execute_reply.started":"2022-12-12T14:35:13.643921Z","shell.execute_reply":"2022-12-12T14:35:34.657267Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset","metadata":{}},{"cell_type":"markdown","source":"## Loading realistic noise","metadata":{}},{"cell_type":"code","source":"df_noise = pd.DataFrame(data=[[f] + list(re.findall('.*/([^/]*)/([^/]*).png', f)[0]) for f in glob.glob(f'{NOISE_DIR}/*/*.png')], columns=['name', 'id', 'detector']).sort_values(['id', 'detector'])\ndf_noise = df_noise.groupby('id').filter(lambda df: len(df) == 2).groupby('id', sort=False).apply(lambda df: df['name'].values).to_frame('files').reset_index()\ndf_noise_train, df_noise_eval = np.array_split(df_noise, [int(len(df_noise) * 0.9)])\ndf_noise","metadata":{"execution":{"iopub.status.busy":"2022-12-12T14:35:34.660222Z","iopub.execute_input":"2022-12-12T14:35:34.660966Z","iopub.status.idle":"2022-12-12T14:35:35.55328Z","shell.execute_reply.started":"2022-12-12T14:35:34.660922Z","shell.execute_reply":"2022-12-12T14:35:35.552257Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if DEBUG:\n    f1, f2 = df_noise.iloc[42].files\n    display(Image.open(f1))\n    display(Image.open(f2))","metadata":{"execution":{"iopub.status.busy":"2022-12-12T14:35:35.556089Z","iopub.execute_input":"2022-12-12T14:35:35.55709Z","iopub.status.idle":"2022-12-12T14:35:35.58552Z","shell.execute_reply.started":"2022-12-12T14:35:35.557049Z","shell.execute_reply":"2022-12-12T14:35:35.5836Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Loading pure signal","metadata":{}},{"cell_type":"code","source":"glob.glob(f'{SIGNAL_DIR}/*')\nSIGNAL_DIR","metadata":{"execution":{"iopub.status.busy":"2022-12-12T14:35:35.586965Z","iopub.execute_input":"2022-12-12T14:35:35.587667Z","iopub.status.idle":"2022-12-12T14:35:35.645003Z","shell.execute_reply.started":"2022-12-12T14:35:35.587629Z","shell.execute_reply":"2022-12-12T14:35:35.644008Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_signal = pd.DataFrame(data=[[f] + list(re.findall('.*/(.*)_(.*).png', f)[0]) for f in glob.glob(f'{SIGNAL_DIR}/*')], columns=['name', 'id', 'detector']).sort_values(['id', 'detector'])\ndf_signal = df_signal.groupby('id').filter(lambda df: len(df) == 2).groupby('id', sort=False).apply(lambda df: df['name'].values).to_frame('files').reset_index()\ndf_signal_train, df_signal_eval = np.array_split(df_signal, [int(len(df_signal) * 0.9)])\ndf_signal","metadata":{"execution":{"iopub.status.busy":"2022-12-12T14:35:35.646314Z","iopub.execute_input":"2022-12-12T14:35:35.646778Z","iopub.status.idle":"2022-12-12T14:35:36.268218Z","shell.execute_reply.started":"2022-12-12T14:35:35.646709Z","shell.execute_reply":"2022-12-12T14:35:36.267119Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if DEBUG:\n    f1, f2 = df_signal.iloc[42].files\n    display(Image.open(f1))\n    display(Image.open(f2))","metadata":{"execution":{"iopub.status.busy":"2022-12-12T14:35:36.270046Z","iopub.execute_input":"2022-12-12T14:35:36.270454Z","iopub.status.idle":"2022-12-12T14:35:36.291368Z","shell.execute_reply.started":"2022-12-12T14:35:36.270416Z","shell.execute_reply":"2022-12-12T14:35:36.290377Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## RealisticNoiseDataset","metadata":{}},{"cell_type":"code","source":"from collections import defaultdict\nimport re\n\ndef get_transforms():\n    return torchvision.transforms.Compose([\n            torchvision.transforms.ToTensor(),\n            torchvision.transforms.Normalize(mean=0.5, std=0.1)\n        ])\n\nclass RealisticNoiseDataset(Dataset):\n    def __init__(self, size, df_noise, df_signal, positive_rate=POSITIVE_RATE, is_train=False) -> None:\n        self.df_noise = df_noise\n        self.df_signal = df_signal\n        self.positive_rate = positive_rate\n        self.size = size\n        self.transforms = get_transforms()\n        self.is_train = is_train\n\n\n    def gen_sample(self, signal, noise, signal_strength):\n        # print(signal, noise)\n        noise = np.array(Image.open(noise))\n        # print(np.mean(noise.flatten() / 255), np.std(noise.flatten()/ 255))\n        if signal:\n            signal = np.array(Image.open(signal))\n            noise = noise + signal_strength * signal\n\n        if self.is_train and GAUSSIAN_NOISE > 0:\n            noise = noise + np.random.randn(*noise.shape) * GAUSSIAN_NOISE \n\n        noise = np.clip(noise, 0, 255).astype(np.uint8)\n        return self.transforms(noise)\n\n\n    def __getitem__(self, index):\n        noise_files = self.df_noise.sample().files.values[0]\n        \n        sig_files = [None, None]\n        label = 0\n        if np.random.random() < self.positive_rate:\n            sig_files = self.df_signal.sample().files.values[0]\n            label = 1\n        signal_strength = np.random.uniform(SIGNAL_LOW, SIGNAL_HIGH)                    \n        return np.concatenate([self.gen_sample(sig, noise, signal_strength) for sig, noise in zip(sig_files, noise_files)], axis=0), label, signal_strength\n\n\n    def __len__(self):\n        return self.size\n\n\n\nds_eval = RealisticNoiseDataset(\n    len(df_signal_eval), \n    df_noise_eval,\n    df_signal_eval\n)\ndl_eval = torch.utils.data.DataLoader(ds_eval, batch_size=BATCH_SIZE, num_workers=os.cpu_count(), pin_memory=True)\n\nif DEBUG:\n    for i in range(4):\n        X, y, ss = ds_eval[i]\n        plt.figure()\n        plt.suptitle(f'{y}, {ss:.02f}')\n        plt.subplot(121).imshow(X[0], cmap='gray')\n        plt.subplot(122).imshow(X[1], cmap='gray')","metadata":{"execution":{"iopub.status.busy":"2022-12-12T14:35:36.29283Z","iopub.execute_input":"2022-12-12T14:35:36.293368Z","iopub.status.idle":"2022-12-12T14:35:37.85228Z","shell.execute_reply.started":"2022-12-12T14:35:36.293333Z","shell.execute_reply":"2022-12-12T14:35:37.851346Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Evaluation","metadata":{}},{"cell_type":"code","source":"def evaluate(model, dl_eval, return_X=False):\n    with torch.no_grad():\n        model.eval()\n        pred = []\n        target = []\n        ss = []\n        signal_strength = []\n        Xs = []\n        for X, y, ss in tqdm(dl_eval, desc='Eval'):\n            pred.append(model(X.to(DEVICE)).cpu().squeeze())\n            target.append(y)\n            signal_strength.append(ss)\n            if return_X:\n                Xs.append(X)\n        pred = torch.concat(pred)\n        target = torch.concat(target)\n        loss = torch.nn.functional.binary_cross_entropy_with_logits(pred, target.float(), reduction='none').median().item() # Avoiding outlier loss with median\n        pred = torch.sigmoid(pred)\n        ret = [roc_auc_score(target, pred), loss, pred, target, torch.concat(signal_strength).numpy()]\n        if return_X:\n            ret.append(torch.concat(Xs).numpy())\n        return ret\n    \nif DEBUG:\n    auc, loss, pred, target, ss, X = evaluate(timm.create_model('inception_v4', pretrained=True, num_classes=1, in_chans=2).to(DEVICE), dl_eval, return_X=True)\n    print(auc, loss)\n    plt.title('(Untrained) distribution of predictions')\n    _ = plt.hist(pred, bins=100)","metadata":{"execution":{"iopub.status.busy":"2022-12-12T14:35:37.853628Z","iopub.execute_input":"2022-12-12T14:35:37.854158Z","iopub.status.idle":"2022-12-12T14:35:47.283206Z","shell.execute_reply.started":"2022-12-12T14:35:37.854123Z","shell.execute_reply":"2022-12-12T14:35:47.28221Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training","metadata":{}},{"cell_type":"code","source":"def train(model, dl_train, epoch, run, optim, scheduler):\n    for step, (X, y, ss) in enumerate(tqdm(dl_train, desc='Train')):\n        pred = model(X.to(DEVICE)).squeeze()\n        loss = torch.nn.functional.binary_cross_entropy_with_logits(pred, y.float().to(DEVICE))\n\n        optim.zero_grad()\n        loss.backward()\n        norm = torch.nn.utils.clip_grad_norm_(model.parameters(), MAX_GRAD_NORM)\n        optim.step()\n        if scheduler:\n            scheduler.step()\n        if run:\n            run.log({\n                'step': step,\n                'loss': loss.item(),\n                'lr': scheduler.get_last_lr()[0] if scheduler else LR,\n                'grad_norm': norm,\n                'epoch': epoch,\n                'logit': pred.mean().item()\n            })\n","metadata":{"execution":{"iopub.status.busy":"2022-12-12T14:35:47.287821Z","iopub.execute_input":"2022-12-12T14:35:47.288772Z","iopub.status.idle":"2022-12-12T14:35:47.296788Z","shell.execute_reply.started":"2022-12-12T14:35:47.28872Z","shell.execute_reply":"2022-12-12T14:35:47.295724Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.model_selection import KFold","metadata":{"execution":{"iopub.status.busy":"2022-12-12T14:35:47.298258Z","iopub.execute_input":"2022-12-12T14:35:47.298679Z","iopub.status.idle":"2022-12-12T14:35:47.309111Z","shell.execute_reply.started":"2022-12-12T14:35:47.298644Z","shell.execute_reply":"2022-12-12T14:35:47.308091Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_dl(fold):\n    kfold = KFold(N_FOLDS, shuffle=True, random_state=42)\n    df_noise_train, df_noise_eval = None, None\n    for f, (train_idx, eval_idx) in enumerate(kfold.split(df_noise)):\n        if f == fold:\n            df_noise_train = df_noise.loc[train_idx]\n            df_noise_eval = df_noise.loc[eval_idx]\n\n    df_signal_train, df_signal_eval = None, None\n    for f, (train_idx, eval_idx) in enumerate(kfold.split(df_signal)):\n        if f == fold:\n            df_signal_train = df_signal.loc[train_idx]\n            df_signal_eval = df_signal.loc[eval_idx]\n\n    ds_train = RealisticNoiseDataset(\n        len(df_signal_train), \n        df_noise_train,\n        df_signal_train,\n        is_train=True\n    )\n\n    ds_eval = RealisticNoiseDataset(\n        len(df_signal_eval), \n        df_noise_eval,\n        df_signal_eval\n    )\n\n    dl_train = torch.utils.data.DataLoader(ds_train, batch_size=BATCH_SIZE, num_workers=os.cpu_count(), pin_memory=True)\n    dl_eval = torch.utils.data.DataLoader(ds_eval, batch_size=BATCH_SIZE, num_workers=os.cpu_count(), pin_memory=True)\n    return dl_train, dl_eval\n\nif DEBUG:\n    dl_train, dl_eval = get_dl(0)\n    for v in dl_train:\n        print(v)\n        break","metadata":{"execution":{"iopub.status.busy":"2022-12-12T14:35:47.310687Z","iopub.execute_input":"2022-12-12T14:35:47.311121Z","iopub.status.idle":"2022-12-12T14:35:48.735923Z","shell.execute_reply.started":"2022-12-12T14:35:47.311086Z","shell.execute_reply":"2022-12-12T14:35:48.734518Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def run_training(run=None, fold=0):\n    dl_train, dl_eval = get_dl(fold)\n\n    model = timm.create_model(MODEL, pretrained=True, num_classes=1, in_chans=2, drop_rate=DROPOUT).to(DEVICE)\n    optim = torch.optim.Adam(model.parameters(), lr=LR)\n    scheduler = None\n    if ONE_CYCLE:\n        scheduler = torch.optim.lr_scheduler.OneCycleLR(optimizer=optim, max_lr=LR, total_steps=len(dl_train) * EPOCHS, pct_start=ONE_CYCLE_PCT_START)\n\n\n    max_auc = 0\n    for epoch in range(EPOCHS):\n        train(model, dl_train, epoch, run, optim, scheduler)\n        auc, loss = evaluate(model, dl_eval)[:2]\n        if auc > max_auc:\n            !mkdir -p models\n            torch.save(model.state_dict(), f'models/model-f{fold}.tph')\n            max_auc = auc\n        if run:\n            run.log({\n                'val_loss': loss,\n                'val_auc': auc,\n                'val_max_auc': max_auc\n            })\n    return max_auc\n\nif TRAIN:\n    for fold in FOLDS:\n        with wandb.init(project='g2net', name=f'{WANDB_RUN}-f{fold}', group=WANDB_RUN) as run:\n            run_training(run, fold=fold)\n","metadata":{"execution":{"iopub.status.busy":"2022-12-12T14:35:48.737812Z","iopub.execute_input":"2022-12-12T14:35:48.738761Z","iopub.status.idle":"2022-12-12T15:16:49.315012Z","shell.execute_reply.started":"2022-12-12T14:35:48.738714Z","shell.execute_reply":"2022-12-12T15:16:49.314082Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%wandb -h 1000 vslaykovsky/g2net","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Optuna","metadata":{}},{"cell_type":"code","source":"import optuna\nfrom optuna.integration.wandb import WeightsAndBiasesCallback\nfrom tqdm.notebook import tqdm\n\n#OPTUNA = True\nif OPTUNA:\n\n    wandb_kwargs = {\"project\": \"g2net-optuna\"}\n    wandbc = WeightsAndBiasesCallback(wandb_kwargs=wandb_kwargs, as_multirun=True)\n\n    @wandbc.track_in_wandb()\n    def objective(trial: optuna.Trial):\n        LR = trial.suggest_float('LR', 0.0001, 0.005, log=True)\n        DROPOUT = trial.suggest_categorical('DROPOUT', [0.05, 0.1,0.2, 0.25,0.3,0.35,0.4])\n        MAX_GRAD_NORM = trial.suggest_float('MAX_GRAD_NORM', 1, 20)\n        EPOCHS = int(trial.suggest_float('EPOCHS', 1, 10, step=1.))\n        GAUSSIAN_NOISE = trial.suggest_categorical('GAUSSIAN_NOISE', [0., 1.])\n        ONE_CYCLE_PCT_START= trial.suggest_categorical('ONE_CYCLE_PCT_START', [0., 0.2])\n        MODEL = trial.suggest_categorical('MODEL', ['resnext50_32x4d', 'efficientnetv2_rw_s', 'seresnext50_32x4d', 'inception_v4'])\n        ONE_CYCLE = trial.suggest_categorical('ONE_CYCLE', [True, False])\n        return run_training(wandb)\n\n\n    study = optuna.create_study(direction='maximize')\n    study.optimize(objective, n_trials=300, callbacks=[wandbc])","metadata":{"execution":{"iopub.status.busy":"2022-12-12T15:16:49.330217Z","iopub.execute_input":"2022-12-12T15:16:49.33073Z","iopub.status.idle":"2022-12-12T15:16:49.539752Z","shell.execute_reply.started":"2022-12-12T15:16:49.330692Z","shell.execute_reply":"2022-12-12T15:16:49.53882Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%wandb vslaykovsky/g2net-optuna","metadata":{"execution":{"iopub.status.busy":"2022-12-12T15:16:49.541016Z","iopub.execute_input":"2022-12-12T15:16:49.541381Z","iopub.status.idle":"2022-12-12T15:16:49.550045Z","shell.execute_reply.started":"2022-12-12T15:16:49.541346Z","shell.execute_reply":"2022-12-12T15:16:49.549077Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Eval set: distributions","metadata":{}},{"cell_type":"code","source":"def load_model(fold=0):\n    model = timm.create_model(MODEL, pretrained=False, num_classes=1, in_chans=2).to(DEVICE)\n    model.load_state_dict(torch.load(f'models/model-f{fold}.tph'))\n    model.eval()\n    return model\n    \nmodel = load_model()","metadata":{"execution":{"iopub.status.busy":"2022-12-12T15:16:49.551499Z","iopub.execute_input":"2022-12-12T15:16:49.552082Z","iopub.status.idle":"2022-12-12T15:16:50.389483Z","shell.execute_reply.started":"2022-12-12T15:16:49.552037Z","shell.execute_reply":"2022-12-12T15:16:50.388499Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preds, targets, signal_strength, X_eval = evaluate(model, dl_eval, return_X=True)[2:]\nplt.title('Eval: distribution of predictions')\n_ = plt.hist(preds, bins=100)","metadata":{"execution":{"iopub.status.busy":"2022-12-12T15:16:50.390946Z","iopub.execute_input":"2022-12-12T15:16:50.391318Z","iopub.status.idle":"2022-12-12T15:17:07.351813Z","shell.execute_reply.started":"2022-12-12T15:16:50.391282Z","shell.execute_reply":"2022-12-12T15:17:07.350746Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ss_buckets = np.linspace(SIGNAL_LOW, SIGNAL_HIGH, 9)\n\nscores = []\nfor ss1, ss2 in zip(ss_buckets[:-1], ss_buckets[1:]):\n    bucket = np.where((ss1 < signal_strength) & (signal_strength <= ss2))\n    score = roc_auc_score(targets[bucket], preds[bucket])\n    scores.append(score)\n\nplt.title('Eval: AUC by signal strength')\nplt.plot(ss_buckets[:-1], scores)","metadata":{"execution":{"iopub.status.busy":"2022-12-12T15:17:07.35625Z","iopub.execute_input":"2022-12-12T15:17:07.358974Z","iopub.status.idle":"2022-12-12T15:17:07.616021Z","shell.execute_reply.started":"2022-12-12T15:17:07.358898Z","shell.execute_reply":"2022-12-12T15:17:07.615073Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Eval set: positive prediction examples","metadata":{}},{"cell_type":"code","source":"idx = preds > 0.8\nfor X, y, pred, ss in list(zip(X_eval[idx], targets[idx], preds[idx], signal_strength[idx]))[:4]:\n    plt.figure()\n    plt.suptitle(f'Target: {y}; prediction: {pred:.02f}; signal_strength: {ss:.02f}')\n    plt.subplot(121).imshow(X[0], cmap='gray')\n    plt.subplot(122).imshow(X[1], cmap='gray')","metadata":{"execution":{"iopub.status.busy":"2022-12-12T15:17:07.620509Z","iopub.execute_input":"2022-12-12T15:17:07.622696Z","iopub.status.idle":"2022-12-12T15:17:09.797776Z","shell.execute_reply.started":"2022-12-12T15:17:07.622658Z","shell.execute_reply":"2022-12-12T15:17:09.796716Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Eval set: negative prediction examples","metadata":{}},{"cell_type":"code","source":"idx = preds < 0.2\nfor X, y, pred, ss in list(zip(X_eval[idx], targets[idx], preds[idx], signal_strength[idx]))[:10]:\n    plt.figure()\n    plt.suptitle(f'Target: {y}; prediction: {pred:.02f}; signal_strength: {ss:.02f}')\n    plt.subplot(121).imshow(X[0], cmap='gray')\n    plt.subplot(122).imshow(X[1], cmap='gray')","metadata":{"execution":{"iopub.status.busy":"2022-12-12T15:17:09.799397Z","iopub.execute_input":"2022-12-12T15:17:09.799963Z","iopub.status.idle":"2022-12-12T15:17:13.974779Z","shell.execute_reply.started":"2022-12-12T15:17:09.799927Z","shell.execute_reply":"2022-12-12T15:17:13.973706Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Prediction","metadata":{}},{"cell_type":"markdown","source":"## Generate test PNGs","metadata":{}},{"cell_type":"code","source":"import pandas as pd\ndf_test = pd.read_csv('/kaggle/input/g2net-winning-strategy-with-external-data/test.csv').query('is_generated_noise == False')\ndf_test","metadata":{"execution":{"iopub.status.busy":"2022-12-12T15:17:13.976558Z","iopub.execute_input":"2022-12-12T15:17:13.976944Z","iopub.status.idle":"2022-12-12T15:17:14.015332Z","shell.execute_reply.started":"2022-12-12T15:17:13.976877Z","shell.execute_reply":"2022-12-12T15:17:14.014248Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pathlib\nimport h5py\n\n# Utility to read hdf5 file\ndef read_data(file):\n    file = pathlib.Path(file)\n    with h5py.File(file, \"r\") as f:\n        filename = file.stem\n        f = f[filename]\n        h1 = f[\"H1\"]\n        l1 = f[\"L1\"]\n        freq_hz = list(f[\"frequency_Hz\"])\n        \n        h1_stft = h1[\"SFTs\"][()]\n        h1_timestamp = h1[\"timestamps_GPS\"][()]\n        # H2 data\n        l1_stft = l1[\"SFTs\"][()]\n        l1_timestamp = l1[\"timestamps_GPS\"][()]\n        \n        return [h1_stft, h1_timestamp],            [l1_stft, l1_timestamp], np.array(freq_hz)\n        \nif DEBUG:\n    [h1_sft, h1_ts], [l1_sft, l1_ts], freq = read_data('/kaggle/input/g2net-detecting-continuous-gravitational-waves/test/00054c878.hdf5')","metadata":{"execution":{"iopub.status.busy":"2022-12-12T15:17:14.017127Z","iopub.execute_input":"2022-12-12T15:17:14.01752Z","iopub.status.idle":"2022-12-12T15:17:14.342119Z","shell.execute_reply.started":"2022-12-12T15:17:14.017483Z","shell.execute_reply":"2022-12-12T15:17:14.341155Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def sft_to_img(sft, ts, buckets=256):\n    bucket_size = (ts.max() - ts.min()) // buckets\n    idx = np.searchsorted(ts, [ts[0] + bucket_size * i for i in range(buckets)])\n    # 1. ASD\n    sft = np.absolute(sft)    \n    # 2. shrink\n    global_noise_amp = np.mean(sft, axis=1)\n    img = np.stack([\n        np.mean(i, axis=1) if i.shape[1] > 0 else global_noise_amp for i in np.array_split(sft, idx[1:], axis=1) \n    ])    \n    ts_noise = img.mean(axis=1)\n    # 3. Normalize\n    mean, std = np.mean(img), np.std(img.astype(np.float64))\n    # print(mean, std)    \n\n    img = img - mean\n    img = img / std / 4.9  # different \n    img *= 128 \n    img += 128\n    # img = img * (255 / np.max(img))\n    img = np.clip(img, 0, 255).astype(np.uint8)\n    # print(np.min(img), np.max(img), np.mean(img), np.std(img))\n    return img.T, ts_noise\n\n\nif DEBUG:\n    img, ts_noise = sft_to_img(h1_sft, h1_ts)\n    plt.figure()\n    plt.imshow(img, cmap='gray')\n    plt.figure()\n    plt.plot(ts_noise)","metadata":{"execution":{"iopub.status.busy":"2022-12-12T15:17:14.343558Z","iopub.execute_input":"2022-12-12T15:17:14.34403Z","iopub.status.idle":"2022-12-12T15:17:14.732423Z","shell.execute_reply.started":"2022-12-12T15:17:14.343992Z","shell.execute_reply":"2022-12-12T15:17:14.731506Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import cv2\n\n!mkdir -p data/test_png\nfor idx, (i, r) in enumerate(tqdm(df_test.iterrows(), total=len(df_test))):\n    if not os.path.exists(f'data/test_png/{r.id}_h1.png'):\n        [h1_sft, h1_ts], [l1_sft, l1_ts], freq = read_data(f'/kaggle/input/g2net-detecting-continuous-gravitational-waves/test/{r.id}.hdf5')\n        h1 = sft_to_img(h1_sft, h1_ts)[0]\n        l1 = sft_to_img(l1_sft, l1_ts)[0]\n        cv2.imwrite(f'data/test_png/{r.id}_h1.png', h1)\n        cv2.imwrite(f'data/test_png/{r.id}_l1.png', l1)    \n    if DEBUG and idx < 5:\n        plt.subplot(121).imshow(Image.open(f'data/test_png/{r.id}_h1.png'), cmap='gray')\n        plt.subplot(122).imshow(Image.open(f'data/test_png/{r.id}_l1.png'), cmap='gray')\n        plt.show()","metadata":{"execution":{"iopub.status.busy":"2022-12-12T15:17:14.733778Z","iopub.execute_input":"2022-12-12T15:17:14.734242Z","iopub.status.idle":"2022-12-12T15:23:51.174849Z","shell.execute_reply.started":"2022-12-12T15:17:14.734205Z","shell.execute_reply":"2022-12-12T15:23:51.173358Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Run Prediction","metadata":{}},{"cell_type":"code","source":"import pandas as pd\ndf_test = pd.read_csv('/kaggle/input/g2net-winning-strategy-with-external-data/test.csv').query('is_generated_noise == False')\ndf_test","metadata":{"execution":{"iopub.status.busy":"2022-12-12T15:23:51.17727Z","iopub.execute_input":"2022-12-12T15:23:51.178051Z","iopub.status.idle":"2022-12-12T15:23:51.211692Z","shell.execute_reply.started":"2022-12-12T15:23:51.178005Z","shell.execute_reply":"2022-12-12T15:23:51.210228Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class G2NETDataset(Dataset):\n    def __init__(self, df, image_path) -> None:\n        self.df = df\n        self.image_path = image_path\n        self.transforms = get_transforms()\n\n    def __getitem__(self, index):\n        id = self.df.iloc[index].id\n        img = np.concatenate([self.transforms(Image.open(f'{self.image_path}/{id}_h1.png')),\n            self.transforms(Image.open(f'{self.image_path}/{id}_l1.png'))])\n        return img\n\n    def __len__(self):\n        return len(self.df)\n\n\n\nds_test = G2NETDataset(df_test, 'data/test_png/')\n\nif DEBUG:\n    X = ds_test[42]\n    plt.subplot(121).imshow(X[0], cmap='gray')\n    plt.subplot(122).imshow(X[1], cmap='gray')","metadata":{"execution":{"iopub.status.busy":"2022-12-12T15:23:51.213709Z","iopub.execute_input":"2022-12-12T15:23:51.214143Z","iopub.status.idle":"2022-12-12T15:23:51.632343Z","shell.execute_reply.started":"2022-12-12T15:23:51.214104Z","shell.execute_reply":"2022-12-12T15:23:51.631246Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\ndl_test = torch.utils.data.DataLoader(ds_test, batch_size=BATCH_SIZE, shuffle=False, num_workers=os.cpu_count())\n\n\nwith torch.no_grad():\n    fold_preds = []\n    for fold in tqdm(FOLDS, desc='Fold prediction'):\n        gc.collect()\n        torch.cuda.empty_cache()\n        \n        model = load_model(fold)\n        preds = torch.concat([torch.sigmoid(model(X.to(DEVICE))).cpu() for X in tqdm(dl_test)], dim=0).numpy()\n        fold_preds.append(preds)\n\npreds = np.stack(fold_preds).squeeze().mean(axis=0)","metadata":{"execution":{"iopub.status.busy":"2022-12-12T15:23:51.640273Z","iopub.execute_input":"2022-12-12T15:23:51.640614Z","iopub.status.idle":"2022-12-12T15:24:51.880427Z","shell.execute_reply.started":"2022-12-12T15:23:51.640582Z","shell.execute_reply":"2022-12-12T15:24:51.879166Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_test['target'] = preds\ndf_test","metadata":{"execution":{"iopub.status.busy":"2022-12-12T15:24:51.882523Z","iopub.execute_input":"2022-12-12T15:24:51.883226Z","iopub.status.idle":"2022-12-12T15:24:51.903457Z","shell.execute_reply.started":"2022-12-12T15:24:51.883159Z","shell.execute_reply":"2022-12-12T15:24:51.902256Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_test.target.plot.hist()","metadata":{"execution":{"iopub.status.busy":"2022-12-12T15:24:51.905671Z","iopub.execute_input":"2022-12-12T15:24:51.906131Z","iopub.status.idle":"2022-12-12T15:24:52.156283Z","shell.execute_reply.started":"2022-12-12T15:24:51.906084Z","shell.execute_reply":"2022-12-12T15:24:52.155176Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(20, 10))\nplt.suptitle('Top positive predictions')\nplt.tight_layout()\nfor idx, id in enumerate(df_test.sort_values('target').tail(10)['id']):\n    plt.subplot(2, 5, idx + 1).imshow(cv2.imread(f'data/test_png/{id}_h1.png'))","metadata":{"execution":{"iopub.status.busy":"2022-12-12T15:24:52.158018Z","iopub.execute_input":"2022-12-12T15:24:52.158598Z","iopub.status.idle":"2022-12-12T15:24:54.058562Z","shell.execute_reply.started":"2022-12-12T15:24:52.158554Z","shell.execute_reply":"2022-12-12T15:24:54.05728Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(20, 10))\nplt.suptitle('Top negative predictions')\nplt.tight_layout()\nfor idx, id in enumerate(df_test.sort_values('target').head(10)['id']):\n    plt.subplot(2, 5, idx + 1).imshow(cv2.imread(f'data/test_png/{id}_h1.png'))","metadata":{"execution":{"iopub.status.busy":"2022-12-12T15:24:54.060489Z","iopub.execute_input":"2022-12-12T15:24:54.060874Z","iopub.status.idle":"2022-12-12T15:24:55.82868Z","shell.execute_reply.started":"2022-12-12T15:24:54.060838Z","shell.execute_reply":"2022-12-12T15:24:55.826908Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Submission + ensemble","metadata":{}},{"cell_type":"code","source":"# df_test = df_test.set_index('id')\ndf_sub = pd.read_csv('/kaggle/input/ensemble-for-public/submission.csv').set_index('id')\ndf_sub.loc[df_test.id, 'target'] =  (0.05*df_test.set_index('id').target + 0.95*df_sub.loc[df_test.id, 'target'])\ndf_sub","metadata":{"execution":{"iopub.status.busy":"2022-12-12T15:24:55.830099Z","iopub.execute_input":"2022-12-12T15:24:55.831068Z","iopub.status.idle":"2022-12-12T15:24:55.896467Z","shell.execute_reply.started":"2022-12-12T15:24:55.831028Z","shell.execute_reply":"2022-12-12T15:24:55.895267Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_sub.to_csv('submission.csv')\n\n!rm -rf data wandb","metadata":{"execution":{"iopub.status.busy":"2022-12-12T15:24:55.897935Z","iopub.execute_input":"2022-12-12T15:24:55.899061Z","iopub.status.idle":"2022-12-12T15:24:57.435532Z","shell.execute_reply.started":"2022-12-12T15:24:55.899012Z","shell.execute_reply":"2022-12-12T15:24:57.434094Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Smash that like button and subscribe for more jaw-dropping notebooks!**","metadata":{}}]}