{"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":"import os\n\nIMG_WIDTH = 512\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\n    if IMG_WIDTH == 512:\n        NOISE_DIR = '/kaggle/input/g2net-realistic-simulation-of-test-noise/data/realistic_noise/images/'\n        SIGNAL_DIR = '/kaggle/input/g2net-generating-pure-signal/data/pure_signal'\n    elif IMG_WIDTH == 256:\n        NOISE_DIR = '/kaggle/input/realistic-noise-256/data/realistic_noise/images/'\n        SIGNAL_DIR = '/kaggle/input/g2net-pure-signal/pure_signal'    \n\n    TEST_CSV = '/kaggle/input/g2net-winning-strategy-with-external-data/test.csv'\n    ROOT = '/kaggle/input/g2net-detecting-continuous-gravitational-waves/'\n    BEST_PUBLIC_SUB_CSV = '/kaggle/input/iknowblendingisnotagoodpracticeonkaggle/submission.csv'\n    IS_KAGGLE = True\n    !pip install -q timm\nexcept:\n    print('Running locally')\n    NOISE_DIR = 'data/realistic_noise/images/'\n    SIGNAL_DIR = 'data/pure_signal/'\n    os.environ['WANDB_API_KEY'] = 'adc8abc0714ba20c3a534b907b9d6beec640f847'\n    TEST_CSV = 'test.csv'\n    ROOT = '/mnt/g2net/'\n    BEST_PUBLIC_SUB_CSV = 'data/prior_submission.csv'\n    IS_KAGGLE = False\n    \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 import tqdm\nimport gc\nimport wandb\nimport pandas as pd\nimport re\n\nplt.rcParams[\"figure.figsize\"] = [15, 7]\n\nBATCH_SIZE = 32\nDEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\n\n    \nPOSITIVE_RATE = 2/3\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# Default\nLR = 0.001\nDROPOUT = 0.\nMAX_GRAD_NORM = 10\nEPOCHS = 1\nGAUSSIAN_NOISE = 2.\nONE_CYCLE_PCT_START=0.1\nMODEL = 'inception_v4'\nONE_CYCLE = True\n\n# Optimal (Optuna)\nDROPOUT=0.25\nEPOCHS=3\nGAUSSIAN_NOISE=2\nLR=0.0006\nMAX_GRAD_NORM=1.5\nMODEL=\"inception_v4\"\nONE_CYCLE=True\nONE_CYCLE_PCT_START=0.1\n\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-11T13:43:17.437654Z","iopub.execute_input":"2022-12-11T13:43:17.438219Z","iopub.status.idle":"2022-12-11T13:43:28.155933Z","shell.execute_reply.started":"2022-12-11T13:43:17.438173Z","shell.execute_reply":"2022-12-11T13:43:28.154749Z"},"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-11T13:35:32.886998Z","iopub.execute_input":"2022-12-11T13:35:32.887627Z","iopub.status.idle":"2022-12-11T13:35:36.052277Z","shell.execute_reply.started":"2022-12-11T13:35:32.887588Z","shell.execute_reply":"2022-12-11T13:35:36.050913Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if DEBUG:\n    f1, f2 = df_noise.iloc[42].files\n    plt.subplot(121).imshow(np.array(Image.open(f1)))\n    plt.subplot(122).imshow(np.array(Image.open(f2)))","metadata":{"execution":{"iopub.status.busy":"2022-12-11T13:35:36.054578Z","iopub.execute_input":"2022-12-11T13:35:36.055227Z","iopub.status.idle":"2022-12-11T13:35:36.401018Z","shell.execute_reply.started":"2022-12-11T13:35:36.055182Z","shell.execute_reply":"2022-12-11T13:35:36.400009Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Loading pure signal","metadata":{}},{"cell_type":"code","source":"if IMG_WIDTH == 256:\n    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'])\nelse:\n    df_signal = pd.DataFrame(data=[[f] + list(re.findall('.*/(.*)/(.*).png', f)[0]) for f in glob.glob(f'{SIGNAL_DIR}/*/*.png')], columns=['name', 'id', 'detector']).sort_values(['id', 'detector'])\n    \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-11T13:35:36.402383Z","iopub.execute_input":"2022-12-11T13:35:36.403373Z","iopub.status.idle":"2022-12-11T13:36:16.121013Z","shell.execute_reply.started":"2022-12-11T13:35:36.403329Z","shell.execute_reply":"2022-12-11T13:36:16.120012Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if DEBUG:\n    f1, f2 = df_signal.iloc[42].files\n    plt.subplot(121).imshow(np.array(Image.open(f1)))\n    plt.subplot(122).imshow(np.array(Image.open(f2)))","metadata":{"execution":{"iopub.status.busy":"2022-12-11T13:36:16.124234Z","iopub.execute_input":"2022-12-11T13:36:16.124641Z","iopub.status.idle":"2022-12-11T13:36:16.441962Z","shell.execute_reply.started":"2022-12-11T13:36:16.124599Z","shell.execute_reply":"2022-12-11T13:36:16.441008Z"},"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#             torchvision.transforms.CenterCrop((224, 224)) # TODO\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-11T13:36:16.443605Z","iopub.execute_input":"2022-12-11T13:36:16.444203Z","iopub.status.idle":"2022-12-11T13:36:17.806497Z","shell.execute_reply.started":"2022-12-11T13:36:16.444155Z","shell.execute_reply":"2022-12-11T13:36:17.805523Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Evaluation","metadata":{}},{"cell_type":"code","source":"def evaluate(model, dl_eval, return_X=False, max_iter=None):\n    gc.collect()\n    torch.cuda.empty_cache()\n    with torch.no_grad():\n        model.eval()\n        pred = []\n        target = []\n        ss = []\n        signal_strength = []\n        Xs = []\n        for idx, (X, y, ss) in enumerate(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            if max_iter is not None and idx >= max_iter:\n                break\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, max_iter=2)\n    print(auc, loss)\n    plt.title('(Untrained) distribution of predictions')\n    _ = plt.hist(pred, bins=100)","metadata":{"execution":{"iopub.status.busy":"2022-12-11T13:40:57.881424Z","iopub.execute_input":"2022-12-11T13:40:57.881847Z","iopub.status.idle":"2022-12-11T13:41:02.974214Z","shell.execute_reply.started":"2022-12-11T13:40:57.881813Z","shell.execute_reply":"2022-12-11T13:41:02.972974Z"},"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-11T13:14:53.900342Z","iopub.execute_input":"2022-12-11T13:14:53.90075Z","iopub.status.idle":"2022-12-11T13:14:53.909332Z","shell.execute_reply.started":"2022-12-11T13:14:53.900705Z","shell.execute_reply":"2022-12-11T13:14:53.908219Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.model_selection import KFold","metadata":{"execution":{"iopub.status.busy":"2022-12-11T13:14:53.911254Z","iopub.execute_input":"2022-12-11T13:14:53.911605Z","iopub.status.idle":"2022-12-11T13:14:53.93066Z","shell.execute_reply.started":"2022-12-11T13:14:53.91157Z","shell.execute_reply":"2022-12-11T13:14:53.929608Z"},"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 X, y, ss in dl_train:\n        print(X.shape, y.shape, ss.shape)\n        break","metadata":{"execution":{"iopub.status.busy":"2022-12-11T13:14:53.931926Z","iopub.execute_input":"2022-12-11T13:14:53.932364Z","iopub.status.idle":"2022-12-11T13:14:56.905264Z","shell.execute_reply.started":"2022-12-11T13:14:53.932329Z","shell.execute_reply":"2022-12-11T13:14:56.903939Z"},"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            os.makedirs('models', exist_ok=True)\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-11T13:14:56.909364Z","iopub.execute_input":"2022-12-11T13:14:56.909789Z","iopub.status.idle":"2022-12-11T13:24:58.050004Z","shell.execute_reply.started":"2022-12-11T13:14:56.909748Z","shell.execute_reply":"2022-12-11T13:24:58.048494Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%wandb -h 1000 vslaykovsky/g2net","metadata":{"execution":{"iopub.status.busy":"2022-12-11T13:24:58.05215Z","iopub.execute_input":"2022-12-11T13:24:58.052541Z","iopub.status.idle":"2022-12-11T13:24:58.064734Z","shell.execute_reply.started":"2022-12-11T13:24:58.052496Z","shell.execute_reply":"2022-12-11T13:24:58.063634Z"},"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        global LR, DROPOUT, MAX_GRAD_NORM, EPOCHS, GAUSSIAN_NOISE, ONE_CYCLE_PCT_START, MODEL, ONE_CYCLE\n\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, 4, step=1.))\n        # GAUSSIAN_NOISE = trial.suggest_float('GAUSSIAN_NOISE', 0, 5)\n        GAUSSIAN_NOISE = trial.suggest_float('GAUSSIAN_NOISE', 0, 5)\n        ONE_CYCLE_PCT_START= trial.suggest_categorical('ONE_CYCLE_PCT_START', [0., 0.1, 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-11T13:24:58.070338Z","iopub.execute_input":"2022-12-11T13:24:58.070652Z","iopub.status.idle":"2022-12-11T13:24:58.277602Z","shell.execute_reply.started":"2022-12-11T13:24:58.070623Z","shell.execute_reply":"2022-12-11T13:24:58.27662Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%wandb vslaykovsky/g2net-optuna","metadata":{"execution":{"iopub.status.busy":"2022-12-11T13:24:58.279161Z","iopub.execute_input":"2022-12-11T13:24:58.279516Z","iopub.status.idle":"2022-12-11T13:24:58.290164Z","shell.execute_reply.started":"2022-12-11T13:24:58.279478Z","shell.execute_reply":"2022-12-11T13:24:58.289014Z"},"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-11T13:47:32.274861Z","iopub.execute_input":"2022-12-11T13:47:32.275626Z","iopub.status.idle":"2022-12-11T13:47:33.149663Z","shell.execute_reply.started":"2022-12-11T13:47:32.275586Z","shell.execute_reply":"2022-12-11T13:47:33.148465Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preds, targets, signal_strength, X_eval = evaluate(model, dl_eval, return_X=True, max_iter=len(dl_eval) // 2)[2:]\nplt.title('Eval: distribution of predictions')\n_ = plt.hist(preds, bins=100)","metadata":{"execution":{"iopub.status.busy":"2022-12-11T13:47:33.1519Z","iopub.execute_input":"2022-12-11T13:47:33.152297Z","iopub.status.idle":"2022-12-11T13:47:51.530714Z","shell.execute_reply.started":"2022-12-11T13:47:33.152259Z","shell.execute_reply":"2022-12-11T13:47:51.527432Z"},"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-11T13:47:55.079883Z","iopub.execute_input":"2022-12-11T13:47:55.081108Z","iopub.status.idle":"2022-12-11T13:47:55.349877Z","shell.execute_reply.started":"2022-12-11T13:47:55.08104Z","shell.execute_reply":"2022-12-11T13:47:55.34888Z"},"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-11T13:47:55.890337Z","iopub.execute_input":"2022-12-11T13:47:55.890736Z","iopub.status.idle":"2022-12-11T13:48:07.078952Z","shell.execute_reply.started":"2022-12-11T13:47:55.8907Z","shell.execute_reply":"2022-12-11T13:48:07.078066Z"},"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]))[: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-11T13:48:07.080815Z","iopub.execute_input":"2022-12-11T13:48:07.08184Z","iopub.status.idle":"2022-12-11T13:48:13.459662Z","shell.execute_reply.started":"2022-12-11T13:48:07.081799Z","shell.execute_reply":"2022-12-11T13:48:13.458806Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del X_eval","metadata":{"execution":{"iopub.status.busy":"2022-12-11T13:48:20.580222Z","iopub.execute_input":"2022-12-11T13:48:20.580594Z","iopub.status.idle":"2022-12-11T13:48:20.680879Z","shell.execute_reply.started":"2022-12-11T13:48:20.580562Z","shell.execute_reply":"2022-12-11T13:48:20.679637Z"},"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(TEST_CSV).query('is_generated_noise == False')\ndf_test","metadata":{"execution":{"iopub.status.busy":"2022-12-11T13:48:24.318939Z","iopub.execute_input":"2022-12-11T13:48:24.31942Z","iopub.status.idle":"2022-12-11T13:48:24.35846Z","shell.execute_reply.started":"2022-12-11T13:48:24.319375Z","shell.execute_reply":"2022-12-11T13:48:24.357174Z"},"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(f'{ROOT}/test/00054c878.hdf5')","metadata":{"execution":{"iopub.status.busy":"2022-12-11T13:48:27.848999Z","iopub.execute_input":"2022-12-11T13:48:27.849433Z","iopub.status.idle":"2022-12-11T13:48:28.292112Z","shell.execute_reply.started":"2022-12-11T13:48:27.849397Z","shell.execute_reply":"2022-12-11T13:48:28.290905Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def sft_to_img(sft, ts, buckets=IMG_WIDTH):\n    \"\"\"\n    The function takes the absolute value of sft, splits it into buckets number of segments, computes the mean of each segment,\n    and normalizes the resulting image to have zero mean and unit variance. \n    It then shifts and scales the values in the image to be within the range 0-255, clips the resulting values to this range, \n    and returns the image and a noise level computed from the mean of the image.\n    \"\"\"\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  # ~ 5 sigma\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-11T13:48:32.271911Z","iopub.execute_input":"2022-12-11T13:48:32.272802Z","iopub.status.idle":"2022-12-11T13:48:32.882984Z","shell.execute_reply.started":"2022-12-11T13:48:32.272762Z","shell.execute_reply":"2022-12-11T13:48:32.881786Z"},"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'{ROOT}/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-11T13:48:40.363904Z","iopub.execute_input":"2022-12-11T13:48:40.364961Z","iopub.status.idle":"2022-12-11T13:48:43.989894Z","shell.execute_reply.started":"2022-12-11T13:48:40.364917Z","shell.execute_reply":"2022-12-11T13:48:43.988642Z"},"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(TEST_CSV).query('is_generated_noise == False')\ndf_test","metadata":{"execution":{"iopub.status.busy":"2022-12-11T13:49:54.655563Z","iopub.execute_input":"2022-12-11T13:49:54.656035Z","iopub.status.idle":"2022-12-11T13:49:54.682495Z","shell.execute_reply.started":"2022-12-11T13:49:54.655984Z","shell.execute_reply":"2022-12-11T13:49:54.681308Z"},"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-11T13:50:09.867755Z","iopub.execute_input":"2022-12-11T13:50:09.868221Z","iopub.status.idle":"2022-12-11T13:50:10.383301Z","shell.execute_reply.started":"2022-12-11T13:50:09.86818Z","shell.execute_reply":"2022-12-11T13:50:10.382373Z"},"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.squeeze())\n\npreds = np.stack(fold_preds).mean(axis=0)","metadata":{"execution":{"iopub.status.busy":"2022-12-11T13:57:13.787635Z","iopub.execute_input":"2022-12-11T13:57:13.788135Z","iopub.status.idle":"2022-12-11T13:57:31.382762Z","shell.execute_reply.started":"2022-12-11T13:57:13.788098Z","shell.execute_reply":"2022-12-11T13:57:31.381412Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preds","metadata":{"execution":{"iopub.status.busy":"2022-12-11T14:00:00.344949Z","iopub.execute_input":"2022-12-11T14:00:00.345401Z","iopub.status.idle":"2022-12-11T14:00:00.353786Z","shell.execute_reply.started":"2022-12-11T14:00:00.345363Z","shell.execute_reply":"2022-12-11T14:00:00.352585Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_test['target'] = preds\ndf_test","metadata":{"execution":{"iopub.status.busy":"2022-12-11T14:00:03.417188Z","iopub.execute_input":"2022-12-11T14:00:03.417626Z","iopub.status.idle":"2022-12-11T14:00:03.436804Z","shell.execute_reply.started":"2022-12-11T14:00:03.417594Z","shell.execute_reply":"2022-12-11T14:00:03.435663Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_test.target.plot.hist()","metadata":{"execution":{"iopub.status.busy":"2022-12-11T14:00:05.51553Z","iopub.execute_input":"2022-12-11T14:00:05.515962Z","iopub.status.idle":"2022-12-11T14:00:05.776899Z","shell.execute_reply.started":"2022-12-11T14:00:05.515928Z","shell.execute_reply":"2022-12-11T14:00:05.775937Z"},"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-11T14:00:13.242927Z","iopub.execute_input":"2022-12-11T14:00:13.244116Z","iopub.status.idle":"2022-12-11T14:00:14.764485Z","shell.execute_reply.started":"2022-12-11T14:00:13.244046Z","shell.execute_reply":"2022-12-11T14:00:14.763493Z"},"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-11T14:00:20.891156Z","iopub.execute_input":"2022-12-11T14:00:20.89154Z","iopub.status.idle":"2022-12-11T14:00:22.394902Z","shell.execute_reply.started":"2022-12-11T14:00:20.891508Z","shell.execute_reply":"2022-12-11T14:00:22.393948Z"},"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(BEST_PUBLIC_SUB_CSV).set_index('id')\ndf_sub.loc[df_test.id, 'target'] =  df_test.set_index('id').target\ndf_sub","metadata":{"execution":{"iopub.status.busy":"2022-12-11T14:00:27.448701Z","iopub.execute_input":"2022-12-11T14:00:27.449107Z","iopub.status.idle":"2022-12-11T14:00:27.483607Z","shell.execute_reply.started":"2022-12-11T14:00:27.449073Z","shell.execute_reply":"2022-12-11T14:00:27.4824Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_sub.to_csv('submission.csv')\n\nif IS_KAGGLE:\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":{}}]}