{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.7.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":37077,"databundleVersionId":4333111,"sourceType":"competition"},{"sourceId":4701755,"sourceType":"datasetVersion","datasetId":2720433},{"sourceId":4704516,"sourceType":"datasetVersion","datasetId":2721684},{"sourceId":4707447,"sourceType":"datasetVersion","datasetId":2722904},{"sourceId":67797978,"sourceType":"kernelVersion"},{"sourceId":112045437,"sourceType":"kernelVersion"},{"sourceId":112547544,"sourceType":"kernelVersion"},{"sourceId":113482597,"sourceType":"kernelVersion"},{"sourceId":113482800,"sourceType":"kernelVersion"}],"dockerImageVersionId":30408,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\ntry:\n    from kaggle_secrets import UserSecretsClient\n    user_secrets = UserSecretsClient()\n    WANDB_API_KEY = user_secrets.get_secret(\"WANDB_API_KEY\")\n    os.environ['WANDB_API_KEY'] = WANDB_API_KEY\n#     NOISE_DIR = '/kaggle/input/g2net-realistic-simulation-of-test-noise/data/realistic_noise/images/'\n    NOISE_DIR = '/kaggle/input/realistic-noise-256/data/realistic_noise/images/'\n    \n#     SIGNAL_DIR = '/kaggle/input/g2net-generating-pure-signal/data/pure_signal'\n    SIGNAL_DIR = '/kaggle/input/g2net-pure-signal/pure_signal'    \n    !pip install -q timm\n    !pip install -q git+https://github.com/PyFstat/PyFstat@python37\nexcept:\n    print('Running locally')\n    NOISE_DIR = 'data/realistic_noise/images/'\n    SIGNAL_DIR = 'data/pure_signal/'\n    os.environ['WANDB_API_KEY'] = 'your_key_here'\n\nfrom timm.data.transforms_factory import create_transform\nimport torchvision\nfrom torch.utils.data import Dataset\nimport glob\nimport numpy as np\nfrom PIL import Image\nimport matplotlib.pyplot as plt\nimport timm\nimport torch\nfrom sklearn.metrics import *\nfrom tqdm 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\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.56 \nDROPOUT = 0.25\nMAX_GRAD_NORM = 1.36\nEPOCHS = 3\nGAUSSIAN_NOISE = 2.\nONE_CYCLE_PCT_START=0.1\nMODEL = 'EfficientNet_Channel_MHA_Attn'\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":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-01-23T15:28:36.39287Z","iopub.execute_input":"2024-01-23T15:28:36.393459Z","iopub.status.idle":"2024-01-23T15:29:47.703614Z","shell.execute_reply.started":"2024-01-23T15:28:36.393408Z","shell.execute_reply":"2024-01-23T15:29:47.702021Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2024-01-23T15:29:47.706485Z","iopub.execute_input":"2024-01-23T15:29:47.70689Z","iopub.status.idle":"2024-01-23T15:29:52.970042Z","shell.execute_reply.started":"2024-01-23T15:29:47.706853Z","shell.execute_reply":"2024-01-23T15:29:52.96868Z"},"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":"2024-01-23T15:29:52.971778Z","iopub.execute_input":"2024-01-23T15:29:52.972151Z","iopub.status.idle":"2024-01-23T15:29:53.026975Z","shell.execute_reply.started":"2024-01-23T15:29:52.972117Z","shell.execute_reply":"2024-01-23T15:29:53.025243Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"glob.glob(f'{SIGNAL_DIR}/*')\nSIGNAL_DIR","metadata":{"execution":{"iopub.status.busy":"2024-01-23T15:29:53.028932Z","iopub.execute_input":"2024-01-23T15:29:53.029489Z","iopub.status.idle":"2024-01-23T15:29:53.45959Z","shell.execute_reply.started":"2024-01-23T15:29:53.02936Z","shell.execute_reply":"2024-01-23T15:29:53.458238Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_signal = pd.DataFrame(data=[[f] + list(re.findall('.*/(.*)_(.*).png', f)[0]) for f in glob.glob(f'{SIGNAL_DIR}/*')], columns=['name', 'id', 'detector']).sort_values(['id', 'detector'])\ndf_signal = df_signal.groupby('id').filter(lambda df: len(df) == 2).groupby('id', sort=False).apply(lambda df: df['name'].values).to_frame('files').reset_index()\ndf_signal_train, df_signal_eval = np.array_split(df_signal, [int(len(df_signal) * 0.9)])\ndf_signal","metadata":{"execution":{"iopub.status.busy":"2024-01-23T15:29:53.464099Z","iopub.execute_input":"2024-01-23T15:29:53.465007Z","iopub.status.idle":"2024-01-23T15:29:54.497755Z","shell.execute_reply.started":"2024-01-23T15:29:53.46494Z","shell.execute_reply":"2024-01-23T15:29:54.496298Z"},"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":"2024-01-23T15:29:54.499828Z","iopub.execute_input":"2024-01-23T15:29:54.500607Z","iopub.status.idle":"2024-01-23T15:29:54.530377Z","shell.execute_reply.started":"2024-01-23T15:29:54.500561Z","shell.execute_reply":"2024-01-23T15:29:54.529289Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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        noise=noise/255.0\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')\n","metadata":{"execution":{"iopub.status.busy":"2024-01-23T15:30:25.719959Z","iopub.execute_input":"2024-01-23T15:30:25.720536Z","iopub.status.idle":"2024-01-23T15:30:27.812982Z","shell.execute_reply.started":"2024-01-23T15:30:25.720487Z","shell.execute_reply":"2024-01-23T15:30:27.811485Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.metrics import roc_auc_score\nimport numpy as np\n\ndef 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_batch = model(X.to(DEVICE)).cpu().squeeze()\n            target_batch = y\n            signal_strength_batch = ss\n\n            # Check for NaN and infinity in pred and target\n            invalid_values = np.isnan(pred_batch.numpy()) | np.isinf(pred_batch.numpy())\n            pred_batch[invalid_values] = 0.0  # Replace NaN and infinity with 0\n\n            target_batch_invalid = np.isnan(target_batch.numpy()) | np.isinf(target_batch.numpy())\n            target_batch[target_batch_invalid] = 0.0  # Replace NaN and infinity with 0\n\n            pred.append(pred_batch)\n            target.append(target_batch)\n            signal_strength.append(signal_strength_batch)\n            if return_X:\n                Xs.append(X)\n\n        pred = torch.cat(pred)\n        target = torch.cat(target)\n        loss = torch.nn.functional.binary_cross_entropy_with_logits(pred, target.float(), reduction='none').median().item()\n        pred = torch.sigmoid(pred)\n\n        ret = [roc_auc_score(target.numpy(), pred.numpy()), loss, pred, target, torch.cat(signal_strength).numpy()]\n        if return_X:\n            ret.append(torch.cat(Xs).numpy())\n        return ret\n","metadata":{"execution":{"iopub.status.busy":"2024-01-23T15:30:29.153762Z","iopub.execute_input":"2024-01-23T15:30:29.154287Z","iopub.status.idle":"2024-01-23T15:30:29.167953Z","shell.execute_reply.started":"2024-01-23T15:30:29.154241Z","shell.execute_reply":"2024-01-23T15:30:29.16669Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\n! pip install efficientnet_pytorch\nfrom efficientnet_pytorch import EfficientNet\n\nclass MultiHeadAttention(nn.Module):\n    def __init__(self, d_model, num_heads):\n        super().__init__()\n        self.d_model = d_model\n        self.num_heads = num_heads\n        \n        assert d_model % num_heads == 0\n        \n        self.d_k = d_model // num_heads\n        \n        self.q_linear = nn.Linear(d_model, d_model)\n        self.v_linear = nn.Linear(d_model, d_model)\n        self.k_linear = nn.Linear(d_model, d_model)\n        \n        self.dropout = nn.Dropout(0.1)\n        self.out_linear = nn.Linear(d_model, d_model)\n        \n    def forward(self, q, k, v, mask=None):\n        bs = q.size(0)\n        \n        # perform linear operation and split into heads\n        q = self.q_linear(q).view(bs, -1, self.num_heads, self.d_k)\n        k = self.k_linear(k).view(bs, -1, self.num_heads, self.d_k)\n        v = self.v_linear(v).view(bs, -1, self.num_heads, self.d_k)\n        \n        # transpose to get dimensions bs * num_heads * sl * d_model\n        q = q.transpose(1,2)\n        k = k.transpose(1,2)\n        v = v.transpose(1,2)\n        \n        # calculate attention using function we will define next\n        scores = self.attention(q, k, v, self.d_k, mask, self.dropout)\n        \n        # concatenate heads and put through final linear layer\n        concat = scores.transpose(1,2).contiguous().view(bs, -1, self.d_model)\n        output = self.out_linear(concat)\n        return output\n    \n    def attention(self, q, k, v, d_k, mask=None, dropout=None):\n        scores = torch.matmul(q, k.transpose(-2, -1)) /  math.sqrt(d_k)\n        if mask is not None:\n            mask = mask.unsqueeze(1)\n            scores = scores.masked_fill(mask == 0, -1e9)\n        scores = nn.functional.softmax(scores, dim=-1)\n        if dropout is not None:\n            scores = dropout(scores)\n        output = torch.matmul(scores, v)\n        return output\n\nclass EfficientNetWithAttention(nn.Module):\n    def __init__(self, num_classes, d_model, num_heads):\n        super().__init__()\n        \n        self.efficient_net = EfficientNet.from_pretrained('efficientnet-b0',in_channels=2)\n        \n        self.pooling = nn.AdaptiveAvgPool2d(1) # add global average pooling layer\n        self.linear1 = nn.Linear(1280, d_model) # add linear layer to adjust input size\n        self.attention = MultiHeadAttention(d_model, num_heads)\n        self.classifier = nn.Linear(d_model, num_classes)\n        \n    def forward(self, x):\n        x = self.efficient_net.extract_features(x)\n        x = self.pooling(x)\n        x = x.view(x.size(0), -1) # flatten\n        x = self.linear1(x) # adjust input size\n        x = self.attention(x, x, x)\n        x = self.classifier(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2024-01-23T15:30:33.501736Z","iopub.execute_input":"2024-01-23T15:30:33.502295Z","iopub.status.idle":"2024-01-23T15:30:47.985558Z","shell.execute_reply.started":"2024-01-23T15:30:33.502242Z","shell.execute_reply":"2024-01-23T15:30:47.983877Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport math\ncustom_model = EfficientNetWithAttention(num_classes=1, d_model=256, num_heads=4)","metadata":{"execution":{"iopub.status.busy":"2024-01-23T15:30:11.547257Z","iopub.execute_input":"2024-01-23T15:30:11.547707Z","iopub.status.idle":"2024-01-23T15:30:12.28899Z","shell.execute_reply.started":"2024-01-23T15:30:11.547667Z","shell.execute_reply":"2024-01-23T15:30:12.287684Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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            })","metadata":{"execution":{"iopub.status.busy":"2024-01-23T15:30:55.104147Z","iopub.execute_input":"2024-01-23T15:30:55.104769Z","iopub.status.idle":"2024-01-23T15:30:55.120126Z","shell.execute_reply.started":"2024-01-23T15:30:55.104717Z","shell.execute_reply":"2024-01-23T15:30:55.118409Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.model_selection import KFold","metadata":{"execution":{"iopub.status.busy":"2024-01-23T15:30:55.804231Z","iopub.execute_input":"2024-01-23T15:30:55.804764Z","iopub.status.idle":"2024-01-23T15:30:55.828422Z","shell.execute_reply.started":"2024-01-23T15:30:55.804722Z","shell.execute_reply":"2024-01-23T15:30:55.827275Z"},"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":"2024-01-23T15:30:56.540986Z","iopub.execute_input":"2024-01-23T15:30:56.54185Z","iopub.status.idle":"2024-01-23T15:30:59.104567Z","shell.execute_reply.started":"2024-01-23T15:30:56.54179Z","shell.execute_reply":"2024-01-23T15:30:59.103199Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dl_train, dl_eval = get_dl(0)\n\n# Assuming dl_train is a list of data\nfor data in dl_train:\n    # Check for missing values in each element of the list\n    if isinstance(data, pd.DataFrame):  # Assuming each element is a Pandas DataFrame\n        missing_values = data.isna().sum()\n        \n        print(\"Missing values in the DataFrame:\")\n        print(missing_values)","metadata":{"execution":{"iopub.status.busy":"2024-01-23T15:38:04.342042Z","iopub.execute_input":"2024-01-23T15:38:04.343109Z","iopub.status.idle":"2024-01-23T15:38:51.521434Z","shell.execute_reply.started":"2024-01-23T15:38:04.343054Z","shell.execute_reply":"2024-01-23T15:38:51.519495Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def run_training(run=None, fold=0):\n    dl_train, dl_eval = get_dl(fold)\n\n    model = custom_model.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)","metadata":{"execution":{"iopub.status.busy":"2024-01-23T15:39:44.394469Z","iopub.execute_input":"2024-01-23T15:39:44.395759Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import optuna\nfrom optuna.integration.wandb import WeightsAndBiasesCallback\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":"2023-04-01T17:21:40.517792Z","iopub.status.idle":"2023-04-01T17:21:40.518315Z","shell.execute_reply.started":"2023-04-01T17:21:40.51806Z","shell.execute_reply":"2023-04-01T17:21:40.518086Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_model(fold=0):\n    model = custom_model.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":"2023-04-01T17:21:40.520166Z","iopub.status.idle":"2023-04-01T17:21:40.520668Z","shell.execute_reply.started":"2023-04-01T17:21:40.520417Z","shell.execute_reply":"2023-04-01T17:21:40.520444Z"},"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":"2023-04-01T17:21:40.52237Z","iopub.status.idle":"2023-04-01T17:21:40.523306Z","shell.execute_reply.started":"2023-04-01T17:21:40.523042Z","shell.execute_reply":"2023-04-01T17:21:40.523075Z"},"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":"2023-04-01T17:21:40.524956Z","iopub.status.idle":"2023-04-01T17:21:40.525927Z","shell.execute_reply.started":"2023-04-01T17:21:40.52563Z","shell.execute_reply":"2023-04-01T17:21:40.525663Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2023-04-01T17:21:40.527565Z","iopub.status.idle":"2023-04-01T17:21:40.528063Z","shell.execute_reply.started":"2023-04-01T17:21:40.527799Z","shell.execute_reply":"2023-04-01T17:21:40.527838Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2023-04-01T17:21:40.529777Z","iopub.status.idle":"2023-04-01T17:21:40.530807Z","shell.execute_reply.started":"2023-04-01T17:21:40.530544Z","shell.execute_reply":"2023-04-01T17:21:40.530571Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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-04-01T17:21:40.532419Z","iopub.status.idle":"2023-04-01T17:21:40.532932Z","shell.execute_reply.started":"2023-04-01T17:21:40.532662Z","shell.execute_reply":"2023-04-01T17:21:40.532688Z"},"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":"2023-04-01T17:21:40.534863Z","iopub.status.idle":"2023-04-01T17:21:40.53542Z","shell.execute_reply.started":"2023-04-01T17:21:40.535125Z","shell.execute_reply":"2023-04-01T17:21:40.53515Z"},"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":"2023-04-01T17:21:40.537214Z","iopub.status.idle":"2023-04-01T17:21:40.537873Z","shell.execute_reply.started":"2023-04-01T17:21:40.537595Z","shell.execute_reply":"2023-04-01T17:21:40.537621Z"},"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":"2023-04-01T17:21:40.539533Z","iopub.status.idle":"2023-04-01T17:21:40.540037Z","shell.execute_reply.started":"2023-04-01T17:21:40.539772Z","shell.execute_reply":"2023-04-01T17:21:40.539796Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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-04-01T17:21:40.541755Z","iopub.status.idle":"2023-04-01T17:21:40.542265Z","shell.execute_reply.started":"2023-04-01T17:21:40.542015Z","shell.execute_reply":"2023-04-01T17:21:40.542041Z"},"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\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":"2023-04-01T17:21:40.543907Z","iopub.status.idle":"2023-04-01T17:21:40.544862Z","shell.execute_reply.started":"2023-04-01T17:21:40.544558Z","shell.execute_reply":"2023-04-01T17:21:40.544591Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dl_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":"2023-04-01T17:21:40.546585Z","iopub.status.idle":"2023-04-01T17:21:40.547088Z","shell.execute_reply.started":"2023-04-01T17:21:40.54684Z","shell.execute_reply":"2023-04-01T17:21:40.546866Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_test['target'] = preds\ndf_test","metadata":{"execution":{"iopub.status.busy":"2023-04-01T17:21:40.548843Z","iopub.status.idle":"2023-04-01T17:21:40.549332Z","shell.execute_reply.started":"2023-04-01T17:21:40.549082Z","shell.execute_reply":"2023-04-01T17:21:40.549108Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_test.target.plot.hist()","metadata":{"execution":{"iopub.status.busy":"2023-04-01T17:21:40.551224Z","iopub.status.idle":"2023-04-01T17:21:40.551727Z","shell.execute_reply.started":"2023-04-01T17:21:40.551471Z","shell.execute_reply":"2023-04-01T17:21:40.551496Z"},"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":"2023-04-01T17:21:40.553475Z","iopub.status.idle":"2023-04-01T17:21:40.553981Z","shell.execute_reply.started":"2023-04-01T17:21:40.553712Z","shell.execute_reply":"2023-04-01T17:21:40.553737Z"},"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":"2023-04-01T17:21:40.5558Z","iopub.status.idle":"2023-04-01T17:21:40.556659Z","shell.execute_reply.started":"2023-04-01T17:21:40.556392Z","shell.execute_reply":"2023-04-01T17:21:40.55642Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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-04-01T17:21:40.557982Z","iopub.status.idle":"2023-04-01T17:21:40.558996Z","shell.execute_reply.started":"2023-04-01T17:21:40.558706Z","shell.execute_reply":"2023-04-01T17:21:40.558738Z"},"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":"2023-04-01T17:21:40.560574Z","iopub.status.idle":"2023-04-01T17:21:40.561077Z","shell.execute_reply.started":"2023-04-01T17:21:40.56081Z","shell.execute_reply":"2023-04-01T17:21:40.560851Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","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","metadata":{"execution":{"iopub.status.busy":"2023-03-17T05:09:41.456378Z","iopub.status.idle":"2023-03-17T05:09:41.45722Z","shell.execute_reply.started":"2023-03-17T05:09:41.456933Z","shell.execute_reply":"2023-03-17T05:09:41.45696Z"}}},{"cell_type":"markdown","source":"df_sub.to_csv('submission.csv',index=False)\n\n!rm -rf data wandb","metadata":{"execution":{"iopub.status.busy":"2023-03-17T05:09:41.458685Z","iopub.status.idle":"2023-03-17T05:09:41.459524Z","shell.execute_reply.started":"2023-03-17T05:09:41.459265Z","shell.execute_reply":"2023-03-17T05:09:41.459291Z"}}},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}