{"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":{"execution":{"iopub.status.busy":"2023-01-02T10:26:30.254853Z","iopub.execute_input":"2023-01-02T10:26:30.258837Z","iopub.status.idle":"2023-01-02T10:26:43.384044Z","shell.execute_reply.started":"2023-01-02T10:26:30.258724Z","shell.execute_reply":"2023-01-02T10:26:43.382239Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install -q git+https://github.com/PyFstat/PyFstat@python37","metadata":{"execution":{"iopub.status.busy":"2023-01-02T10:26:43.388097Z","iopub.execute_input":"2023-01-02T10:26:43.389196Z","iopub.status.idle":"2023-01-02T10:27:22.394099Z","shell.execute_reply.started":"2023-01-02T10:26:43.389144Z","shell.execute_reply":"2023-01-02T10:27:22.392892Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import shutil","metadata":{"execution":{"iopub.status.busy":"2023-01-02T10:27:22.395521Z","iopub.execute_input":"2023-01-02T10:27:22.395929Z","iopub.status.idle":"2023-01-02T10:27:22.40386Z","shell.execute_reply.started":"2023-01-02T10:27:22.395886Z","shell.execute_reply":"2023-01-02T10:27:22.40287Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!mkdir merged_pure_signals","metadata":{"execution":{"iopub.status.busy":"2023-01-02T10:27:22.408332Z","iopub.execute_input":"2023-01-02T10:27:22.408703Z","iopub.status.idle":"2023-01-02T10:27:23.355286Z","shell.execute_reply.started":"2023-01-02T10:27:22.408669Z","shell.execute_reply":"2023-01-02T10:27:23.353644Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\n# sss_file_dir = \"/kaggle/input/pure-signal-sss/data/pure_signal_SSS\"\n# target_path = \"/kaggle/working/merged_pure_signals\"\n# for group in os.listdir(sss_file_dir):\n#     if '.' not in group:\n#         for file in os.listdir(os.path.join(sss_file_dir, group)):\n#             shutil.copyfile(os.path.join(sss_file_dir, group, file), os.path.join(target_path, 'sss' + group + \"_\" + file))\n\nsss_file_dir = '/kaggle/input/g2net-pure-signal/pure_signal'\ntarget_path = \"/kaggle/working/merged_pure_signals\"\n\nfor file in os.listdir(sss_file_dir):\n    shutil.copyfile(os.path.join(sss_file_dir, file), os.path.join(target_path, file))","metadata":{"execution":{"iopub.status.busy":"2023-01-02T10:27:23.357486Z","iopub.execute_input":"2023-01-02T10:27:23.357924Z","iopub.status.idle":"2023-01-02T10:28:47.113976Z","shell.execute_reply.started":"2023-01-02T10:27:23.357878Z","shell.execute_reply":"2023-01-02T10:28:47.112956Z"},"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 = 36\nDEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\n\n\n# try:\nfrom kaggle_secrets import UserSecretsClient\nuser_secrets = UserSecretsClient()\nWANDB_API_KEY = user_secrets.get_secret(\"VVV\")\nos.environ['WANDB_API_KEY'] = WANDB_API_KEY\nNOISE_DIR = '/kaggle/input/realistic-noise-256/data/realistic_noise/images/'\nSIGNAL_DIR = target_path\n\n# except:\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 = False\nOPTUNA = False\nTRAIN = True\nFOLDS = [0, 1, 2, 3, 4, 5]\n# FOLDS = [0] \nN_FOLDS = len(FOLDS)\n\n# class Config:\nLR = 0.00056\nDROPOUT = 0.25\nMAX_GRAD_NORM = 1.36\nEPOCHS = 10\nGAUSSIAN_NOISE = 2.\nONE_CYCLE_PCT_START=0.1\nMODEL = 'tf_efficientnetv2_s_in21k'\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":"2023-01-02T10:28:47.115639Z","iopub.execute_input":"2023-01-02T10:28:47.115998Z","iopub.status.idle":"2023-01-02T10:28:50.677216Z","shell.execute_reply.started":"2023-01-02T10:28:47.11596Z","shell.execute_reply":"2023-01-02T10:28:50.676285Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset","metadata":{}},{"cell_type":"code","source":"ps_percent = 0.02\ndf_test = pd.read_csv('/kaggle/input/g2net-winning-strategy-with-external-data/test.csv').query('is_generated_noise == False')\n# df_test = pd.read_csv(\"/kaggle/input/g2net-detecting-continuous-gravitational-waves/sample_submission.csv\")\npl_labeling = pd.read_csv(\"/kaggle/input/best-sub-for-pl/submission(9).csv\")\npl_labeling2 = pd.read_csv(\"/kaggle/input/pseoudo2/submission(10).csv\")\n\npl_upper_bound = pl_labeling['target'].quantile(1 - ps_percent)\npl_lower_bound = pl_labeling['target'].quantile(ps_percent)\n\npl_labeling['new_target'] = 0.5\npl_labeling.loc[pl_labeling['target'] > pl_upper_bound, ['new_target']] = 1\npl_labeling.loc[pl_labeling['target'] < pl_lower_bound, ['new_target']] = 0\n\npl_data = pl_labeling[pl_labeling['new_target'] != 0.5][['id', 'new_target']].reset_index(drop=True)\nprint(pl_data.shape)\npl_data = pl_data[pl_data['id'].astype(str).isin(df_test['id'].astype(str))].reset_index(drop=True).copy()\n\npl_labeling['target_pseudo'] = pl_labeling2['target']\n\npl_labeling['diff'] = np.abs(pl_labeling['target_pseudo'] - pl_labeling['target'])\n\npl_labeling.sort_values('diff').to_csv(\"pseudo_handcheck.csv\")","metadata":{"execution":{"iopub.status.busy":"2023-01-02T11:27:28.4497Z","iopub.execute_input":"2023-01-02T11:27:28.450063Z","iopub.status.idle":"2023-01-02T11:27:28.543073Z","shell.execute_reply.started":"2023-01-02T11:27:28.450032Z","shell.execute_reply":"2023-01-02T11:27:28.542072Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.plot(pl_labeling.sort_values('diff', ascending=False)['diff'].values)","metadata":{"execution":{"iopub.status.busy":"2023-01-02T11:27:37.29917Z","iopub.execute_input":"2023-01-02T11:27:37.299655Z","iopub.status.idle":"2023-01-02T11:27:37.500356Z","shell.execute_reply.started":"2023-01-02T11:27:37.299615Z","shell.execute_reply":"2023-01-02T11:27:37.49946Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\ndf_test\n\nimport 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')\n\ndef 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)\n\nimport 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":"2023-01-02T11:27:37.875297Z","iopub.execute_input":"2023-01-02T11:27:37.876392Z","iopub.status.idle":"2023-01-02T11:55:15.932467Z","shell.execute_reply.started":"2023-01-02T11:27:37.876351Z","shell.execute_reply":"2023-01-02T11:55:15.931247Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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)\n\n\ndf_noise = df_noise.to_frame('files').reset_index()\ndf_noise_train, df_noise_eval = np.array_split(df_noise, [int(len(df_noise) * 0.9)])\ndf_noise\n\npl_data['files'] = pl_data['id'].apply(lambda x: np.array([f'/kaggle/working/data/test_png/{x}_h1.png', f'/kaggle/working/data/test_png/{x}_l1.png'], ))\n\nnoise_test = pl_data[pl_data['new_target'] < 1.0].copy().reset_index(drop=True)\n\ndf_noise= pd.concat([df_noise, noise_test[['id', 'files']]]).reset_index(drop=True)\n\ndf_noise = df_noise[['id', 'files']].copy()","metadata":{"execution":{"iopub.status.busy":"2023-01-02T10:34:58.996465Z","iopub.execute_input":"2023-01-02T10:34:58.997226Z","iopub.status.idle":"2023-01-02T10:35:02.81674Z","shell.execute_reply.started":"2023-01-02T10:34:58.997179Z","shell.execute_reply":"2023-01-02T10:35:02.815761Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if DEBUG:\n    f1, f2 = df_noise.iloc[1495].files\n    display(Image.open(f1))\n    display(Image.open(f2))","metadata":{"execution":{"iopub.status.busy":"2023-01-02T10:35:02.820774Z","iopub.execute_input":"2023-01-02T10:35:02.821082Z","iopub.status.idle":"2023-01-02T10:35:02.826128Z","shell.execute_reply.started":"2023-01-02T10:35:02.821054Z","shell.execute_reply":"2023-01-02T10:35:02.824694Z"},"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)])\n\nsignal_test = pl_data[pl_data['new_target'] > 0.0].copy().reset_index(drop=True)[['id', 'files']]\n\ndf_signal = pd.concat([df_signal, signal_test]).reset_index(drop=True)","metadata":{"execution":{"iopub.status.busy":"2023-01-02T10:35:02.828347Z","iopub.execute_input":"2023-01-02T10:35:02.829247Z","iopub.status.idle":"2023-01-02T10:35:03.432267Z","shell.execute_reply.started":"2023-01-02T10:35:02.829208Z","shell.execute_reply":"2023-01-02T10:35:03.431317Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if DEBUG:\n    f1, f2 = df_signal.iloc[9628].files\n    display(Image.open(f1))\n    display(Image.open(f2))","metadata":{"execution":{"iopub.status.busy":"2023-01-02T10:35:03.43357Z","iopub.execute_input":"2023-01-02T10:35:03.434548Z","iopub.status.idle":"2023-01-02T10:35:03.439971Z","shell.execute_reply.started":"2023-01-02T10:35:03.434509Z","shell.execute_reply":"2023-01-02T10:35:03.438649Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## RealisticNoiseDataset","metadata":{}},{"cell_type":"code","source":"import numpy as np\n\nfrom 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        noise = np.array(Image.open(noise))\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":"2023-01-02T10:35:03.441789Z","iopub.execute_input":"2023-01-02T10:35:03.442141Z","iopub.status.idle":"2023-01-02T10:35:03.462788Z","shell.execute_reply.started":"2023-01-02T10:35:03.442105Z","shell.execute_reply":"2023-01-02T10:35:03.461673Z"},"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(MODEL, 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":"2023-01-02T10:35:03.46426Z","iopub.execute_input":"2023-01-02T10:35:03.465167Z","iopub.status.idle":"2023-01-02T10:35:03.477083Z","shell.execute_reply.started":"2023-01-02T10:35:03.465138Z","shell.execute_reply":"2023-01-02T10:35:03.476123Z"},"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":"2023-01-02T10:35:03.480079Z","iopub.execute_input":"2023-01-02T10:35:03.480941Z","iopub.status.idle":"2023-01-02T10:35:03.490786Z","shell.execute_reply.started":"2023-01-02T10:35:03.480912Z","shell.execute_reply":"2023-01-02T10:35:03.489811Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.model_selection import KFold","metadata":{"execution":{"iopub.status.busy":"2023-01-02T10:35:03.494105Z","iopub.execute_input":"2023-01-02T10:35:03.494363Z","iopub.status.idle":"2023-01-02T10:35:03.537711Z","shell.execute_reply.started":"2023-01-02T10:35:03.494338Z","shell.execute_reply":"2023-01-02T10:35:03.536793Z"},"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) * 3, \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":"2023-01-02T10:35:03.539328Z","iopub.execute_input":"2023-01-02T10:35:03.539871Z","iopub.status.idle":"2023-01-02T10:35:03.551346Z","shell.execute_reply.started":"2023-01-02T10:35:03.539834Z","shell.execute_reply":"2023-01-02T10:35:03.550336Z"},"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    print(len(dl_train))\n    print(len(dl_eval))\n    \n\n    model = timm.create_model(MODEL, pretrained=True, num_classes=1, in_chans=2, drop_rate=DROPOUT).to(DEVICE)\n    optim = torch.optim.RAdam(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        else:\n            break\n        \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":"2023-01-02T10:35:03.554508Z","iopub.execute_input":"2023-01-02T10:35:03.554797Z","iopub.status.idle":"2023-01-02T10:46:04.768165Z","shell.execute_reply.started":"2023-01-02T10:35:03.554772Z","shell.execute_reply":"2023-01-02T10:46:04.766463Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dl_train, dl_eval = get_dl(4)\nprint(len(dl_train))\nprint(len(dl_eval))\n\n\nmodel = timm.create_model(MODEL, pretrained=True, num_classes=1, in_chans=2, drop_rate=DROPOUT).to(DEVICE)\noptim = torch.optim.RAdam(model.parameters(), lr=LR)\nscheduler = None\nif 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\nmax_auc = 0\nfor 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    else:\n        break","metadata":{"execution":{"iopub.status.busy":"2023-01-01T21:19:34.818961Z","iopub.execute_input":"2023-01-01T21:19:34.819408Z","iopub.status.idle":"2023-01-01T21:19:46.413818Z","shell.execute_reply.started":"2023-01-01T21:19:34.819372Z","shell.execute_reply":"2023-01-01T21:19:46.412136Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%wandb -h 1000 vslaykovsky/g2net","metadata":{"execution":{"iopub.status.busy":"2022-12-31T10:08:48.171135Z","iopub.status.idle":"2022-12-31T10:08:48.171642Z","shell.execute_reply.started":"2022-12-31T10:08:48.171363Z","shell.execute_reply":"2022-12-31T10:08:48.171387Z"},"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-31T10:08:48.173524Z","iopub.status.idle":"2022-12-31T10:08:48.174006Z","shell.execute_reply.started":"2022-12-31T10:08:48.173748Z","shell.execute_reply":"2022-12-31T10:08:48.17377Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%wandb vslaykovsky/g2net-optuna","metadata":{"execution":{"iopub.status.busy":"2022-12-31T10:08:48.175741Z","iopub.status.idle":"2022-12-31T10:08:48.176239Z","shell.execute_reply.started":"2022-12-31T10:08:48.175984Z","shell.execute_reply":"2022-12-31T10:08:48.176007Z"},"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-31T10:08:48.178127Z","iopub.status.idle":"2022-12-31T10:08:48.178591Z","shell.execute_reply.started":"2022-12-31T10:08:48.178348Z","shell.execute_reply":"2022-12-31T10:08:48.17837Z"},"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-31T10:08:48.180412Z","iopub.status.idle":"2022-12-31T10:08:48.180926Z","shell.execute_reply.started":"2022-12-31T10:08:48.180662Z","shell.execute_reply":"2022-12-31T10:08:48.180685Z"},"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-31T10:08:48.182693Z","iopub.status.idle":"2022-12-31T10:08:48.183234Z","shell.execute_reply.started":"2022-12-31T10:08:48.182933Z","shell.execute_reply":"2022-12-31T10:08:48.182957Z"},"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-31T10:08:48.185088Z","iopub.status.idle":"2022-12-31T10:08:48.18556Z","shell.execute_reply.started":"2022-12-31T10:08:48.185318Z","shell.execute_reply":"2022-12-31T10:08:48.18534Z"},"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-31T10:08:48.187595Z","iopub.status.idle":"2022-12-31T10:08:48.188638Z","shell.execute_reply.started":"2022-12-31T10:08:48.188374Z","shell.execute_reply":"2022-12-31T10:08:48.188398Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Prediction","metadata":{}},{"cell_type":"markdown","source":"## Generate test PNGs","metadata":{}},{"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":"2023-01-02T10:54:53.640822Z","iopub.execute_input":"2023-01-02T10:54:53.64119Z","iopub.status.idle":"2023-01-02T10:54:53.662978Z","shell.execute_reply.started":"2023-01-02T10:54:53.641157Z","shell.execute_reply":"2023-01-02T10:54:53.661944Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_test.query('is_generated_noise == False')","metadata":{"execution":{"iopub.status.busy":"2023-01-02T10:56:07.266778Z","iopub.execute_input":"2023-01-02T10:56:07.267149Z","iopub.status.idle":"2023-01-02T10:56:07.283161Z","shell.execute_reply.started":"2023-01-02T10:56:07.267117Z","shell.execute_reply":"2023-01-02T10:56:07.281514Z"},"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-31T10:08:48.195446Z","iopub.status.idle":"2022-12-31T10:08:48.195934Z","shell.execute_reply.started":"2022-12-31T10:08:48.195669Z","shell.execute_reply":"2022-12-31T10:08:48.195692Z"},"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-31T10:08:48.197862Z","iopub.status.idle":"2022-12-31T10:08:48.198335Z","shell.execute_reply.started":"2022-12-31T10:08:48.198086Z","shell.execute_reply":"2022-12-31T10:08:48.198109Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_test['target'] = preds\ndf_test","metadata":{"execution":{"iopub.status.busy":"2022-12-31T10:08:48.200199Z","iopub.status.idle":"2022-12-31T10:08:48.200662Z","shell.execute_reply.started":"2022-12-31T10:08:48.20042Z","shell.execute_reply":"2022-12-31T10:08:48.200442Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_test.target.plot.hist()","metadata":{"execution":{"iopub.status.busy":"2022-12-31T10:08:48.203386Z","iopub.status.idle":"2022-12-31T10:08:48.204423Z","shell.execute_reply.started":"2022-12-31T10:08:48.204167Z","shell.execute_reply":"2022-12-31T10:08:48.204192Z"},"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-31T10:08:48.206001Z","iopub.status.idle":"2022-12-31T10:08:48.206852Z","shell.execute_reply.started":"2022-12-31T10:08:48.20657Z","shell.execute_reply":"2022-12-31T10:08:48.206595Z"},"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-31T10:08:48.208224Z","iopub.status.idle":"2022-12-31T10:08:48.209069Z","shell.execute_reply.started":"2022-12-31T10:08:48.208811Z","shell.execute_reply":"2022-12-31T10:08:48.208839Z"},"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')\ndf_sub.loc[df_test.id, 'target'] =  (df_test.set_index('id').target + df_sub.loc[df_test.id, 'target']) / 2\ndf_sub","metadata":{"execution":{"iopub.status.busy":"2023-01-02T10:57:19.905859Z","iopub.execute_input":"2023-01-02T10:57:19.906563Z","iopub.status.idle":"2023-01-02T10:57:19.972052Z","shell.execute_reply.started":"2023-01-02T10:57:19.906508Z","shell.execute_reply":"2023-01-02T10:57:19.96866Z"},"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-31T10:08:48.213876Z","iopub.status.idle":"2022-12-31T10:08:48.214764Z","shell.execute_reply.started":"2022-12-31T10:08:48.214502Z","shell.execute_reply":"2022-12-31T10:08:48.214526Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Smash that like button and subscribe for more jaw-dropping notebooks!**","metadata":{}}]}