{"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":"!pip install -q timm","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install -q git+https://github.com/PyFstat/PyFstat@python37","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from 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 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\n\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/realistic-noise-256/data/realistic_noise/images/'\n    SIGNAL_DIR = '/kaggle/input/g2net-pure-signal/pure_signal'\n    \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    \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:\nLR = 0.00056\nDROPOUT = 0.25\nMAX_GRAD_NORM = 1.36\nEPOCHS = 3\nGAUSSIAN_NOISE = 2.\nONE_CYCLE_PCT_START=0.1\nMODEL = 'inception_v4'\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-11T00:10:51.489763Z","iopub.execute_input":"2022-12-11T00:10:51.490168Z","iopub.status.idle":"2022-12-11T00:10:51.55496Z","shell.execute_reply.started":"2022-12-11T00:10:51.490109Z","shell.execute_reply":"2022-12-11T00:10:51.553277Z"},"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-10T11:52:23.504457Z","iopub.execute_input":"2022-12-10T11:52:23.504927Z","iopub.status.idle":"2022-12-10T11:52:24.424982Z","shell.execute_reply.started":"2022-12-10T11:52:23.504898Z","shell.execute_reply":"2022-12-10T11:52:24.424082Z"},"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-10T11:52:24.426543Z","iopub.execute_input":"2022-12-10T11:52:24.426912Z","iopub.status.idle":"2022-12-10T11:52:24.461271Z","shell.execute_reply.started":"2022-12-10T11:52:24.426875Z","shell.execute_reply":"2022-12-10T11:52:24.460396Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Loading pure signal","metadata":{}},{"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-10T11:52:24.464414Z","iopub.execute_input":"2022-12-10T11:52:24.465058Z","iopub.status.idle":"2022-12-10T11:52:25.090101Z","shell.execute_reply.started":"2022-12-10T11:52:24.465021Z","shell.execute_reply":"2022-12-10T11:52:25.088904Z"},"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-10T11:52:25.091975Z","iopub.execute_input":"2022-12-10T11:52:25.092425Z","iopub.status.idle":"2022-12-10T11:52:25.114088Z","shell.execute_reply.started":"2022-12-10T11:52:25.092366Z","shell.execute_reply":"2022-12-10T11:52:25.112899Z"},"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-10T11:52:25.115821Z","iopub.execute_input":"2022-12-10T11:52:25.11621Z","iopub.status.idle":"2022-12-10T11:52:26.469372Z","shell.execute_reply.started":"2022-12-10T11:52:25.116173Z","shell.execute_reply":"2022-12-10T11:52:26.468294Z"},"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-10T11:52:26.470965Z","iopub.execute_input":"2022-12-10T11:52:26.471449Z","iopub.status.idle":"2022-12-10T11:52:38.610355Z","shell.execute_reply.started":"2022-12-10T11:52:26.471409Z","shell.execute_reply":"2022-12-10T11:52:38.609308Z"},"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-10T11:52:38.612244Z","iopub.execute_input":"2022-12-10T11:52:38.612581Z","iopub.status.idle":"2022-12-10T11:52:38.620787Z","shell.execute_reply.started":"2022-12-10T11:52:38.612547Z","shell.execute_reply":"2022-12-10T11:52:38.619871Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.model_selection import KFold","metadata":{"execution":{"iopub.status.busy":"2022-12-10T11:52:38.624221Z","iopub.execute_input":"2022-12-10T11:52:38.624867Z","iopub.status.idle":"2022-12-10T11:52:38.649411Z","shell.execute_reply.started":"2022-12-10T11:52:38.624812Z","shell.execute_reply":"2022-12-10T11:52:38.646596Z"},"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-10T11:52:38.652614Z","iopub.execute_input":"2022-12-10T11:52:38.652945Z","iopub.status.idle":"2022-12-10T11:52:40.420479Z","shell.execute_reply.started":"2022-12-10T11:52:38.652915Z","shell.execute_reply":"2022-12-10T11:52:40.419299Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\ndef 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-10T11:52:40.425653Z","iopub.execute_input":"2022-12-10T11:52:40.429046Z","iopub.status.idle":"2022-12-10T11:55:45.193724Z","shell.execute_reply.started":"2022-12-10T11:52:40.428948Z","shell.execute_reply":"2022-12-10T11:55:45.192494Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%wandb -h 1000 vslaykovsky/g2net","metadata":{"execution":{"iopub.status.busy":"2022-12-10T11:55:45.19537Z","iopub.execute_input":"2022-12-10T11:55:45.195677Z","iopub.status.idle":"2022-12-10T11:55:45.205066Z","shell.execute_reply.started":"2022-12-10T11:55:45.195642Z","shell.execute_reply":"2022-12-10T11:55:45.204107Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Optuna","metadata":{}},{"cell_type":"code","source":"import optuna\nfrom optuna.integration.wandb import WeightsAndBiasesCallback\n\n\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., 0.1, 0.25, 0.5])\n        MAX_GRAD_NORM = trial.suggest_float('MAX_GRAD_NORM', 1, 20)\n        EPOCHS = int(trial.suggest_float('EPOCHS', 1, 5, 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.1])\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-10T11:55:45.208999Z","iopub.execute_input":"2022-12-10T11:55:45.209664Z","iopub.status.idle":"2022-12-10T11:55:45.405295Z","shell.execute_reply.started":"2022-12-10T11:55:45.209627Z","shell.execute_reply":"2022-12-10T11:55:45.404283Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%wandb vslaykovsky/g2net-optuna","metadata":{"execution":{"iopub.status.busy":"2022-12-10T11:55:45.406595Z","iopub.execute_input":"2022-12-10T11:55:45.40745Z","iopub.status.idle":"2022-12-10T11:55:45.416463Z","shell.execute_reply.started":"2022-12-10T11:55:45.407411Z","shell.execute_reply":"2022-12-10T11:55:45.415339Z"},"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-10T11:55:45.417959Z","iopub.execute_input":"2022-12-10T11:55:45.418623Z","iopub.status.idle":"2022-12-10T11:55:46.253718Z","shell.execute_reply.started":"2022-12-10T11:55:45.418584Z","shell.execute_reply":"2022-12-10T11:55:46.252642Z"},"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-10T11:55:46.255418Z","iopub.execute_input":"2022-12-10T11:55:46.255808Z","iopub.status.idle":"2022-12-10T11:56:03.606853Z","shell.execute_reply.started":"2022-12-10T11:55:46.25577Z","shell.execute_reply":"2022-12-10T11:56:03.605762Z"},"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-10T11:56:03.6096Z","iopub.execute_input":"2022-12-10T11:56:03.610326Z","iopub.status.idle":"2022-12-10T11:56:03.843723Z","shell.execute_reply.started":"2022-12-10T11:56:03.610283Z","shell.execute_reply":"2022-12-10T11:56:03.842758Z"},"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-10T11:56:03.845369Z","iopub.execute_input":"2022-12-10T11:56:03.845742Z","iopub.status.idle":"2022-12-10T11:56:05.927405Z","shell.execute_reply.started":"2022-12-10T11:56:03.845705Z","shell.execute_reply":"2022-12-10T11:56:05.926084Z"},"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-10T11:56:05.929311Z","iopub.execute_input":"2022-12-10T11:56:05.929695Z","iopub.status.idle":"2022-12-10T11:56:09.47271Z","shell.execute_reply.started":"2022-12-10T11:56:05.929658Z","shell.execute_reply":"2022-12-10T11:56:09.471703Z"},"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-10T11:57:53.980748Z","iopub.execute_input":"2022-12-10T11:57:53.981522Z","iopub.status.idle":"2022-12-10T11:57:54.021393Z","shell.execute_reply.started":"2022-12-10T11:57:53.981483Z","shell.execute_reply":"2022-12-10T11:57:54.020486Z"},"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-10T11:58:19.06464Z","iopub.execute_input":"2022-12-10T11:58:19.065047Z","iopub.status.idle":"2022-12-10T11:58:19.893039Z","shell.execute_reply.started":"2022-12-10T11:58:19.06501Z","shell.execute_reply":"2022-12-10T11:58:19.891989Z"},"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-10T11:58:29.866894Z","iopub.execute_input":"2022-12-10T11:58:29.867586Z","iopub.status.idle":"2022-12-10T11:58:30.248809Z","shell.execute_reply.started":"2022-12-10T11:58:29.86755Z","shell.execute_reply":"2022-12-10T11:58:30.247766Z"},"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-10T11:59:13.134365Z","iopub.execute_input":"2022-12-10T11:59:13.134737Z","iopub.status.idle":"2022-12-10T11:59:39.481676Z","shell.execute_reply.started":"2022-12-10T11:59:13.134705Z","shell.execute_reply":"2022-12-10T11:59:39.4802Z"},"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-10T11:59:57.548618Z","iopub.execute_input":"2022-12-10T11:59:57.549004Z","iopub.status.idle":"2022-12-10T11:59:57.572309Z","shell.execute_reply.started":"2022-12-10T11:59:57.548971Z","shell.execute_reply":"2022-12-10T11:59:57.571202Z"},"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-10T11:59:59.635815Z","iopub.execute_input":"2022-12-10T11:59:59.636214Z","iopub.status.idle":"2022-12-10T11:59:59.972628Z","shell.execute_reply.started":"2022-12-10T11:59:59.63618Z","shell.execute_reply":"2022-12-10T11:59:59.971646Z"},"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-10T12:00:02.101712Z","iopub.execute_input":"2022-12-10T12:00:02.102106Z","iopub.status.idle":"2022-12-10T12:00:05.404089Z","shell.execute_reply.started":"2022-12-10T12:00:02.102071Z","shell.execute_reply":"2022-12-10T12:00:05.39952Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_test['target'] = preds\ndf_test","metadata":{"execution":{"iopub.status.busy":"2022-12-10T12:00:12.129183Z","iopub.execute_input":"2022-12-10T12:00:12.129604Z","iopub.status.idle":"2022-12-10T12:00:12.160049Z","shell.execute_reply.started":"2022-12-10T12:00:12.129563Z","shell.execute_reply":"2022-12-10T12:00:12.158842Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_test.to_csv(\"single_model.csv\")","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_test.target.plot.hist()","metadata":{"execution":{"iopub.status.busy":"2022-12-10T11:56:09.695362Z","iopub.status.idle":"2022-12-10T11:56:09.695853Z","shell.execute_reply.started":"2022-12-10T11:56:09.695591Z","shell.execute_reply":"2022-12-10T11:56:09.695614Z"},"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-10T11:56:09.697515Z","iopub.status.idle":"2022-12-10T11:56:09.698016Z","shell.execute_reply.started":"2022-12-10T11:56:09.697742Z","shell.execute_reply":"2022-12-10T11:56:09.697765Z"},"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-10T11:56:09.700501Z","iopub.status.idle":"2022-12-10T11:56:09.701621Z","shell.execute_reply.started":"2022-12-10T11:56:09.701348Z","shell.execute_reply":"2022-12-10T11:56:09.701373Z"},"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/iknowblendingisnotagoodpracticeonkaggle/submission.csv').set_index('id')\ntry_overfitting = pd.read_csv('/kaggle/input/g2netarchive/submission.csv').set_index('id')\ndf_sub.loc[df_test.id, 'target'] =  0.75*(df_test.set_index('id').target + 0.25*df_sub.loc[df_test.id, 'target'])\ndf_sub.loc[~df_sub.index.isin(df_test.index), 'target'] =  0.4*(try_overfitting.loc[~df_sub.index.isin(df_test.index), 'target']) + 0.6*df_sub.loc[~df_sub.index.isin(df_test.index), 'target']\ndf_sub","metadata":{"execution":{"iopub.status.busy":"2022-12-10T12:00:16.890698Z","iopub.execute_input":"2022-12-10T12:00:16.891099Z","iopub.status.idle":"2022-12-10T12:00:16.938559Z","shell.execute_reply.started":"2022-12-10T12:00:16.891064Z","shell.execute_reply":"2022-12-10T12:00:16.937096Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# From https://www.kaggle.com/code/tanreinama/eliminate-noise-using-signal-similarity\nr = []\nfor i, t in zip(df_sub.index, df_sub.target):\n    if i in [\"308417080\", \"dc2aaaee9\", \"8b180f74f\", \"698567d90\"]:\n        r.append(1.0)\n    else:\n        if t < 0.05:\n            t = 0.0\n        elif t > 0.95:\n            t = 1.0\n        r.append(t)\ndf_sub[\"target\"] = r\ndf_sub.to_csv(\"submission.csv\", index=False)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_sub.to_csv('submission.csv')\n\n!rm -rf data wandb","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Smash that like button and subscribe for more jaw-dropping notebooks!**","metadata":{}}]}