{"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":"code","source":"!pip install -q nnAudio\n!pip install -q ttach","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2021-09-11T07:48:40.241758Z","iopub.execute_input":"2021-09-11T07:48:40.242629Z","iopub.status.idle":"2021-09-11T07:48:58.071273Z","shell.execute_reply.started":"2021-09-11T07:48:40.242484Z","shell.execute_reply":"2021-09-11T07:48:58.070421Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!curl https://raw.githubusercontent.com/pytorch/xla/master/contrib/scripts/env-setup.py -o pytorch-xla-env-setup.py\n!python pytorch-xla-env-setup.py --version 1.8\n#!curl https://raw.githubusercontent.com/pytorch/xla/master/contrib/scripts/env-setup.py -o pytorch-xla-env-setup.py\n#!python pytorch-xla-env-setup.py --version \"nightly\"","metadata":{"execution":{"iopub.status.busy":"2021-09-11T07:48:58.073329Z","iopub.execute_input":"2021-09-11T07:48:58.074304Z","iopub.status.idle":"2021-09-11T07:49:58.080002Z","shell.execute_reply.started":"2021-09-11T07:48:58.074258Z","shell.execute_reply":"2021-09-11T07:49:58.078998Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Loading","metadata":{"papermill":{"duration":0.016769,"end_time":"2021-07-01T14:31:32.675036","exception":false,"start_time":"2021-07-01T14:31:32.658267","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image\nfrom matplotlib import pyplot as plt\nimport seaborn as sns","metadata":{"papermill":{"duration":0.717732,"end_time":"2021-07-01T14:31:33.409839","exception":false,"start_time":"2021-07-01T14:31:32.692107","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2021-09-11T07:49:58.081395Z","iopub.execute_input":"2021-09-11T07:49:58.081661Z","iopub.status.idle":"2021-09-11T07:49:58.880558Z","shell.execute_reply.started":"2021-09-11T07:49:58.081629Z","shell.execute_reply":"2021-09-11T07:49:58.879755Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = pd.read_csv('../input/g2net-gravitational-wave-detection/training_labels.csv')\ntest = pd.read_csv('../input/g2net-gravitational-wave-detection/sample_submission.csv')\n\ndef get_train_file_path(image_id):\n    return \"../input/g2net-gravitational-wave-detection/train/{}/{}/{}/{}.npy\".format(\n        image_id[0], image_id[1], image_id[2], image_id)\n\ndef get_test_file_path(image_id):\n    return \"../input/g2net-gravitational-wave-detection/test/{}/{}/{}/{}.npy\".format(\n        image_id[0], image_id[1], image_id[2], image_id)\n\ntrain['file_path'] = train['id'].apply(get_train_file_path)\ntest['file_path'] = test['id'].apply(get_test_file_path)\n\ndisplay(train.head())\ndisplay(test.head())","metadata":{"papermill":{"duration":1.251215,"end_time":"2021-07-01T14:31:34.678928","exception":false,"start_time":"2021-07-01T14:31:33.427713","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2021-09-11T07:49:58.883403Z","iopub.execute_input":"2021-09-11T07:49:58.883777Z","iopub.status.idle":"2021-09-11T07:50:00.311866Z","shell.execute_reply.started":"2021-09-11T07:49:58.883731Z","shell.execute_reply":"2021-09-11T07:50:00.310971Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Quick EDA","metadata":{"papermill":{"duration":0.018004,"end_time":"2021-07-01T14:31:34.717088","exception":false,"start_time":"2021-07-01T14:31:34.699084","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import torch\nfrom nnAudio.Spectrogram import CQT1992v2\n\ndef apply_qtransform(waves, transform=CQT1992v2(sr=2048, fmin=20, fmax=1024, hop_length=64)):\n    waves = np.hstack(waves)\n    waves = waves / np.max(waves)\n    waves = torch.from_numpy(waves).float()\n    image = transform(waves)\n    return image\n\nfor i in range(5):\n    waves = np.load(train.loc[i, 'file_path'])\n    image = apply_qtransform(waves)\n    target = train.loc[i, 'target']\n    plt.imshow(image[0])\n    plt.title(f\"target: {target}\")\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2021-09-11T07:50:00.31345Z","iopub.execute_input":"2021-09-11T07:50:00.31376Z","iopub.status.idle":"2021-09-11T07:50:02.26435Z","shell.execute_reply.started":"2021-09-11T07:50:00.31372Z","shell.execute_reply":"2021-09-11T07:50:02.263662Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train['target'].hist()","metadata":{"papermill":{"duration":0.193214,"end_time":"2021-07-01T14:31:37.746016","exception":false,"start_time":"2021-07-01T14:31:37.552802","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2021-09-11T07:50:02.265305Z","iopub.execute_input":"2021-09-11T07:50:02.266153Z","iopub.status.idle":"2021-09-11T07:50:02.523207Z","shell.execute_reply.started":"2021-09-11T07:50:02.266114Z","shell.execute_reply":"2021-09-11T07:50:02.522497Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Directory settings","metadata":{"papermill":{"duration":0.030055,"end_time":"2021-07-01T14:31:37.805204","exception":false,"start_time":"2021-07-01T14:31:37.775149","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# ====================================================\n# Directory settings\n# ====================================================\nimport os\n\nOUTPUT_DIR = './'\nif not os.path.exists(OUTPUT_DIR):\n    os.makedirs(OUTPUT_DIR)","metadata":{"papermill":{"duration":0.036354,"end_time":"2021-07-01T14:31:37.870775","exception":false,"start_time":"2021-07-01T14:31:37.834421","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2021-09-11T07:50:02.524456Z","iopub.execute_input":"2021-09-11T07:50:02.525035Z","iopub.status.idle":"2021-09-11T07:50:02.52948Z","shell.execute_reply.started":"2021-09-11T07:50:02.525003Z","shell.execute_reply":"2021-09-11T07:50:02.528689Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# CFG","metadata":{"papermill":{"duration":0.028932,"end_time":"2021-07-01T14:31:37.928007","exception":false,"start_time":"2021-07-01T14:31:37.899075","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# ====================================================\n# CFG\n# ====================================================\nclass CFG:\n    apex=True\n    debug=True\n    num_workers=4\n    model_name='tf_efficientnet_b4_ns'\n    scheduler='CosineAnnealingLR' # ['ReduceLROnPlateau', 'CosineAnnealingLR', 'CosineAnnealingWarmRestarts']\n    epochs=8\n    #factor=0.2 # ReduceLROnPlateau\n    #patience=4 # ReduceLROnPlateau\n    #eps=1e-6 # ReduceLROnPlateau\n    T_max=3 # CosineAnnealingLR\n    #T_0=3 # CosineAnnealingWarmRestarts\n    lr=1e-4\n    min_lr=1e-6\n    batch_size=48\n    weight_decay=1e-6\n    gradient_accumulation_steps=1\n    max_grad_norm=1000\n    qtransform_params={\"sr\": 2048, \"fmin\": 20, \"fmax\": 1024, \"hop_length\": 32, \"bins_per_octave\": 8, \"verbose\": False}\n    seed=2021\n    target_size=1\n    target_col='target'\n    n_fold=5\n    trn_fold=[0] # [0, 1, 2, 3, 4]\n    train=True\n    \nif CFG.debug:\n    CFG.epochs = 3\n    train = train.sample(n=10000, random_state=CFG.seed).reset_index(drop=True)","metadata":{"papermill":{"duration":0.181532,"end_time":"2021-07-01T14:31:38.138409","exception":false,"start_time":"2021-07-01T14:31:37.956877","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2021-09-11T07:50:02.530463Z","iopub.execute_input":"2021-09-11T07:50:02.530703Z","iopub.status.idle":"2021-09-11T07:50:02.6742Z","shell.execute_reply.started":"2021-09-11T07:50:02.530675Z","shell.execute_reply":"2021-09-11T07:50:02.673378Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Library","metadata":{"papermill":{"duration":0.028374,"end_time":"2021-07-01T14:31:38.19586","exception":false,"start_time":"2021-07-01T14:31:38.167486","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# ====================================================\n# Library\n# ====================================================\nimport sys\nsys.path.append('../input/pytorch-image-models/pytorch-image-models-master')\n\nimport os\nimport math\nimport time\nimport random\nimport shutil\nfrom pathlib import Path\nfrom contextlib import contextmanager\nfrom collections import defaultdict, Counter\n\nimport scipy as sp\nimport numpy as np\nimport pandas as pd\n\nfrom sklearn import preprocessing\nfrom sklearn.metrics import roc_auc_score\nfrom sklearn.model_selection import StratifiedKFold, GroupKFold, KFold\n\nfrom tqdm.auto import tqdm\nfrom functools import partial\n\nimport cv2\nfrom PIL import Image\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.optim import Adam, SGD\nimport torchvision.models as models\nfrom torch.nn.parameter import Parameter\nfrom torch.utils.data import DataLoader, Dataset\nfrom torch.optim.lr_scheduler import CosineAnnealingWarmRestarts, CosineAnnealingLR, ReduceLROnPlateau\n\nimport torch_xla\nimport torch_xla.debug.metrics as met\nimport torch_xla.distributed.data_parallel as dp\nimport torch_xla.distributed.parallel_loader as pl\nimport torch_xla.utils.utils as xu\nimport torch_xla.core.xla_model as xm\nimport torch_xla.utils.serialization as xser\nimport torch_xla.distributed.xla_multiprocessing as xmp\nimport torch_xla.test.test_utils as test_utils\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom albumentations import ImageOnlyTransform\n\nimport timm\n\nfrom torch.cuda.amp import autocast, GradScaler\n\nimport warnings\nwarnings.filterwarnings('ignore')","metadata":{"papermill":{"duration":3.270545,"end_time":"2021-07-01T14:31:41.495669","exception":false,"start_time":"2021-07-01T14:31:38.225124","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2021-09-11T07:50:02.676491Z","iopub.execute_input":"2021-09-11T07:50:02.676818Z","iopub.status.idle":"2021-09-11T07:50:05.79839Z","shell.execute_reply.started":"2021-09-11T07:50:02.676776Z","shell.execute_reply":"2021-09-11T07:50:05.797438Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Utils","metadata":{"papermill":{"duration":0.028157,"end_time":"2021-07-01T14:31:41.552625","exception":false,"start_time":"2021-07-01T14:31:41.524468","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# ====================================================\n# Utils\n# ====================================================\ndef get_score(y_true, y_pred):\n    score = roc_auc_score(y_true, y_pred)\n    return score\n\n\ndef init_logger(log_file=OUTPUT_DIR+'train.log'):\n    from logging import getLogger, INFO, FileHandler,  Formatter,  StreamHandler\n    logger = getLogger(__name__)\n    logger.setLevel(INFO)\n    handler1 = StreamHandler()\n    handler1.setFormatter(Formatter(\"%(message)s\"))\n    handler2 = FileHandler(filename=log_file)\n    handler2.setFormatter(Formatter(\"%(message)s\"))\n    logger.addHandler(handler1)\n    logger.addHandler(handler2)\n    return logger\n\nLOGGER = init_logger()\n\n\ndef seed_torch(seed=42):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n\nseed_torch(seed=CFG.seed)","metadata":{"papermill":{"duration":0.042029,"end_time":"2021-07-01T14:31:41.623242","exception":false,"start_time":"2021-07-01T14:31:41.581213","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2021-09-11T07:50:05.800672Z","iopub.execute_input":"2021-09-11T07:50:05.801018Z","iopub.status.idle":"2021-09-11T07:50:05.811245Z","shell.execute_reply.started":"2021-09-11T07:50:05.80099Z","shell.execute_reply":"2021-09-11T07:50:05.810575Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# CV split","metadata":{"papermill":{"duration":0.02877,"end_time":"2021-07-01T14:31:41.680818","exception":false,"start_time":"2021-07-01T14:31:41.652048","status":"completed"},"tags":[]}},{"cell_type":"code","source":"Fold = StratifiedKFold(n_splits=CFG.n_fold, shuffle=True, random_state=CFG.seed)\nfor n, (train_index, val_index) in enumerate(Fold.split(train, train[CFG.target_col])):\n    train.loc[val_index, 'fold'] = int(n)\ntrain['fold'] = train['fold'].astype(int)\ndisplay(train.groupby(['fold', 'target']).size())","metadata":{"papermill":{"duration":0.060375,"end_time":"2021-07-01T14:31:41.769944","exception":false,"start_time":"2021-07-01T14:31:41.709569","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2021-09-11T07:50:05.81211Z","iopub.execute_input":"2021-09-11T07:50:05.812667Z","iopub.status.idle":"2021-09-11T07:50:05.848216Z","shell.execute_reply.started":"2021-09-11T07:50:05.81262Z","shell.execute_reply":"2021-09-11T07:50:05.847438Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{"papermill":{"duration":0.028894,"end_time":"2021-07-01T14:31:41.827575","exception":false,"start_time":"2021-07-01T14:31:41.798681","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# ====================================================\n# Dataset\n# ====================================================\nclass TrainDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df\n        self.file_names = df['file_path'].values\n        self.labels = df[CFG.target_col].values\n        self.wave_transform = CQT1992v2(**CFG.qtransform_params)\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def apply_qtransform(self, waves, transform):\n        waves = np.hstack(waves)\n        waves = waves / np.max(waves)\n        waves = torch.from_numpy(waves).float()\n        image = transform(waves)\n        return image\n\n    def __getitem__(self, idx):\n        file_path = self.file_names[idx]\n        waves = np.load(file_path)\n        image = self.apply_qtransform(waves, self.wave_transform)\n        image = image.squeeze().numpy()\n        if self.transform:\n            image = self.transform(image=image)['image']\n        label = torch.tensor(self.labels[idx]).float()\n        return image, label","metadata":{"papermill":{"duration":0.040385,"end_time":"2021-07-01T14:31:41.897587","exception":false,"start_time":"2021-07-01T14:31:41.857202","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2021-09-11T07:50:05.849331Z","iopub.execute_input":"2021-09-11T07:50:05.84967Z","iopub.status.idle":"2021-09-11T07:50:05.859691Z","shell.execute_reply.started":"2021-09-11T07:50:05.849638Z","shell.execute_reply":"2021-09-11T07:50:05.85836Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Transforms","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Transforms\n# ====================================================\ndef get_transforms(*, data):\n    \n    if data == 'train':\n        return A.Compose([\n            ToTensorV2(),\n        ])\n\n    elif data == 'valid':\n        return A.Compose([\n            ToTensorV2(),\n        ])","metadata":{"execution":{"iopub.status.busy":"2021-09-11T07:50:05.861135Z","iopub.execute_input":"2021-09-11T07:50:05.861656Z","iopub.status.idle":"2021-09-11T07:50:05.876066Z","shell.execute_reply.started":"2021-09-11T07:50:05.861611Z","shell.execute_reply":"2021-09-11T07:50:05.875163Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = TrainDataset(train, transform=get_transforms(data='train'))\n\nfor i in range(5):\n    plt.figure(figsize=(16,12))\n    image, label = train_dataset[i]\n    plt.imshow(image[0])\n    plt.title(f'label: {label}')\n    plt.show() ","metadata":{"papermill":{"duration":1.037231,"end_time":"2021-07-01T14:31:42.96351","exception":false,"start_time":"2021-07-01T14:31:41.926279","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2021-09-11T07:50:05.877298Z","iopub.execute_input":"2021-09-11T07:50:05.877522Z","iopub.status.idle":"2021-09-11T07:50:07.125568Z","shell.execute_reply.started":"2021-09-11T07:50:05.877498Z","shell.execute_reply":"2021-09-11T07:50:07.124985Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# MODEL","metadata":{"papermill":{"duration":0.03649,"end_time":"2021-07-01T14:31:43.035743","exception":false,"start_time":"2021-07-01T14:31:42.999253","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# ====================================================\n# MODEL\n# ====================================================\nclass CustomModel(nn.Module):\n    def __init__(self, cfg, pretrained=False):\n        super().__init__()\n        self.cfg = cfg\n        self.model = timm.create_model(self.cfg.model_name, pretrained=pretrained, in_chans=1)\n        self.n_features = self.model.classifier.in_features\n        self.model.classifier = nn.Linear(self.n_features, self.cfg.target_size)\n\n    def forward(self, x):\n        output = self.model(x)\n        return output\n\nMX = xmp.MpModelWrapper(CustomModel(CFG, pretrained=True))\n","metadata":{"papermill":{"duration":0.044023,"end_time":"2021-07-01T14:31:43.114443","exception":false,"start_time":"2021-07-01T14:31:43.07042","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2021-09-11T07:50:07.12646Z","iopub.execute_input":"2021-09-11T07:50:07.127366Z","iopub.status.idle":"2021-09-11T07:50:08.69809Z","shell.execute_reply.started":"2021-09-11T07:50:07.127306Z","shell.execute_reply":"2021-09-11T07:50:08.697167Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Helper functions","metadata":{"papermill":{"duration":0.034387,"end_time":"2021-07-01T14:31:43.183231","exception":false,"start_time":"2021-07-01T14:31:43.148844","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# ====================================================\n# Helper functions\n# ====================================================\nclass AverageMeter(object):\n    \"\"\"Computes and stores the average and current value\"\"\"\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.val = 0\n        self.avg = 0\n        self.sum = 0\n        self.count = 0\n\n    def update(self, val, n=1):\n        self.val = val\n        self.sum += val * n\n        self.count += n\n        self.avg = self.sum / self.count\n\n\ndef asMinutes(s):\n    m = math.floor(s / 60)\n    s -= m * 60\n    return '%dm %ds' % (m, s)\n\n\ndef timeSince(since, percent):\n    now = time.time()\n    s = now - since\n    es = s / (percent)\n    rs = es - s\n    return '%s (remain %s)' % (asMinutes(s), asMinutes(rs))\n\ndef loss_fn(outputs, targets):\n        return nn.BCEWithLogitsLoss()(outputs, targets.view(-1, 1))\n\ndef train_fn(fold, train_loader, model, criterion, optimizer, epoch, scheduler, device):\n    if CFG.apex:\n        scaler = GradScaler()\n    batch_time = AverageMeter()\n    data_time = AverageMeter()\n    losses = AverageMeter()\n    scores = AverageMeter()\n    # switch to train mode\n    model.train()\n    start = end = time.time()\n    global_step = 0\n    for step, (images, labels) in enumerate(train_loader):\n        # measure data loading time\n        data_time.update(time.time() - end)\n        images = images.to(device)\n        labels = labels.to(device)\n        batch_size = labels.size(0)\n        if CFG.apex:\n            with autocast():\n                y_preds = model(images)\n                loss = loss_fn(y_preds, labels)\n        else:\n            y_preds = model(images)\n            loss = criterion(y_preds.view(-1), labels)\n        # record loss\n        losses.update(loss.item(), batch_size)\n        if CFG.gradient_accumulation_steps > 1:\n            loss = loss / CFG.gradient_accumulation_steps\n        #if CFG.apex:\n        #    scaler.scale(loss).backward()\n        #else:\n        #    loss.backward()\n        loss.backward()\n        grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), CFG.max_grad_norm)\n        if (step + 1) % CFG.gradient_accumulation_steps == 0:\n            if CFG.apex:\n                # scaler.step(optimizer)\n                xm.optimizer_step(optimizer)\n                scaler.update()\n            else:\n                optimizer.step()\n            optimizer.zero_grad()\n            global_step += 1\n        # measure elapsed time\n        batch_time.update(time.time() - end)\n        end = time.time()\n    return losses.avg\n\n\ndef valid_fn(valid_loader, model, criterion, device):\n    batch_time = AverageMeter()\n    data_time = AverageMeter()\n    losses = AverageMeter()\n    scores = AverageMeter()\n    # switch to evaluation mode\n    model.eval()\n    preds = []\n    start = end = time.time()\n    valid_labels = []\n    for step, (images, labels) in enumerate(valid_loader):\n        # measure data loading time\n        data_time.update(time.time() - end)\n        images = images.to(device)\n        labels = labels.to(device)\n        batch_size = labels.size(0)\n        # compute loss\n        with torch.no_grad():\n            y_preds = model(images)\n        xm.mark_step()\n        loss = loss_fn(y_preds, labels)\n        losses.update(loss.item(), batch_size)\n        # record accuracy\n        preds.append(y_preds.sigmoid().to('cpu').numpy())\n        valid_labels.append(labels.to('cpu').numpy())\n        if CFG.gradient_accumulation_steps > 1:\n            loss = loss / CFG.gradient_accumulation_steps\n        # measure elapsed time\n        batch_time.update(time.time() - end)\n        end = time.time()\n    preds = np.concatenate(preds)\n    valid_labels = np.concatenate(valid_labels)\n    score = get_score(valid_labels, preds)\n    return losses.avg, preds, score","metadata":{"papermill":{"duration":0.190966,"end_time":"2021-07-01T14:31:43.408934","exception":false,"start_time":"2021-07-01T14:31:43.217968","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2021-09-11T07:50:08.700787Z","iopub.execute_input":"2021-09-11T07:50:08.701189Z","iopub.status.idle":"2021-09-11T07:50:08.728091Z","shell.execute_reply.started":"2021-09-11T07:50:08.701156Z","shell.execute_reply":"2021-09-11T07:50:08.727476Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MX = xmp.MpModelWrapper(CustomModel(CFG, pretrained=True))","metadata":{"execution":{"iopub.status.busy":"2021-09-11T07:50:08.729135Z","iopub.execute_input":"2021-09-11T07:50:08.729559Z","iopub.status.idle":"2021-09-11T07:50:09.20148Z","shell.execute_reply.started":"2021-09-11T07:50:08.72953Z","shell.execute_reply":"2021-09-11T07:50:09.200772Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def epoch_update_gamma(y_true,y_pred, epoch=-1,delta=1):\n        \"\"\"\n        Calculate gamma from last epoch's targets and predictions.\n        Gamma is updated at the end of each epoch.\n        y_true: `Tensor`. Targets (labels).  Float either 0.0 or 1.0 .\n        y_pred: `Tensor` . Predictions.\n        \"\"\"\n        DELTA = delta+1\n        SUB_SAMPLE_SIZE = 2000.0\n        pos = y_pred[y_true==1]\n        neg = y_pred[y_true==0] # yo pytorch, no boolean tensors or operators?  Wassap?\n        # subsample the training set for performance\n        cap_pos = pos.shape[0]\n        cap_neg = neg.shape[0]\n        pos = pos[torch.rand_like(pos) < SUB_SAMPLE_SIZE/cap_pos]\n        neg = neg[torch.rand_like(neg) < SUB_SAMPLE_SIZE/cap_neg]\n        ln_pos = pos.shape[0]\n        ln_neg = neg.shape[0]\n        pos_expand = pos.view(-1,1).expand(-1,ln_neg).reshape(-1)\n        neg_expand = neg.repeat(ln_pos)\n        diff = neg_expand - pos_expand\n        ln_All = diff.shape[0]\n        Lp = diff[diff>0] # because we're taking positive diffs, we got pos and neg flipped.\n        ln_Lp = Lp.shape[0]-1\n        diff_neg = -1.0 * diff[diff<0]\n        diff_neg = diff_neg.sort()[0]\n        ln_neg = diff_neg.shape[0]-1\n        ln_neg = max([ln_neg, 0])\n        left_wing = int(ln_Lp*DELTA)\n        left_wing = max([0,left_wing])\n        left_wing = min([ln_neg,left_wing])\n        default_gamma=torch.tensor(0.2, dtype=torch.float).cuda()\n        if diff_neg.shape[0] > 0 :\n           gamma = diff_neg[left_wing]\n        else:\n           gamma = default_gamma # default=torch.tensor(0.2, dtype=torch.float).cuda() #zoink\n        L1 = diff[diff>-1.0*gamma]\n        ln_L1 = L1.shape[0]\n        if epoch > -1 :\n            return gamma\n        else :\n            return default_gamma\n\n\n\ndef roc_star_loss( _y_true, y_pred, gamma, _epoch_true, epoch_pred):\n        \"\"\"\n        Nearly direct loss function for AUC.\n        See article,\n        C. Reiss, \"Roc-star : An objective function for ROC-AUC that actually works.\"\n        https://github.com/iridiumblue/articles/blob/master/roc_star.md\n            _y_true: `Tensor`. Targets (labels).  Float either 0.0 or 1.0 .\n            y_pred: `Tensor` . Predictions.\n            gamma  : `Float` Gamma, as derived from last epoch.\n            _epoch_true: `Tensor`.  Targets (labels) from last epoch.\n            epoch_pred : `Tensor`.  Predicions from last epoch.\n        \"\"\"\n        #convert labels to boolean\n        y_true = (_y_true>=0.50)\n        epoch_true = (_epoch_true>=0.50)\n\n        # if batch is either all true or false return small random stub value.\n        if torch.sum(y_true)==0 or torch.sum(y_true) == y_true.shape[0]: return torch.sum(y_pred)*1e-8\n\n        pos = y_pred[y_true]\n        neg = y_pred[~y_true]\n\n        epoch_pos = epoch_pred[epoch_true]\n        epoch_neg = epoch_pred[~epoch_true]\n\n        # Take random subsamples of the training set, both positive and negative.\n        max_pos = 1000 # Max number of positive training samples\n        max_neg = 1000 # Max number of positive training samples\n        cap_pos = epoch_pos.shape[0]\n        cap_neg = epoch_neg.shape[0]\n        epoch_pos = epoch_pos[torch.rand_like(epoch_pos) < max_pos/cap_pos]\n        epoch_neg = epoch_neg[torch.rand_like(epoch_neg) < max_neg/cap_pos]\n\n        ln_pos = pos.shape[0]\n        ln_neg = neg.shape[0]\n\n        # sum positive batch elements agaionst (subsampled) negative elements\n        if ln_pos>0 :\n            pos_expand = pos.view(-1,1).expand(-1,epoch_neg.shape[0]).reshape(-1)\n            neg_expand = epoch_neg.repeat(ln_pos)\n\n            diff2 = neg_expand - pos_expand + gamma\n            l2 = diff2[diff2>0]\n            m2 = l2 * l2\n            len2 = l2.shape[0]\n        else:\n            m2 = torch.tensor([0], dtype=torch.float).cuda()\n            len2 = 0\n\n        # Similarly, compare negative batch elements against (subsampled) positive elements\n        if ln_neg>0 :\n            pos_expand = epoch_pos.view(-1,1).expand(-1, ln_neg).reshape(-1)\n            neg_expand = neg.repeat(epoch_pos.shape[0])\n\n            diff3 = neg_expand - pos_expand + gamma\n            l3 = diff3[diff3>0]\n            m3 = l3*l3\n            len3 = l3.shape[0]\n        else:\n            m3 = torch.tensor([0], dtype=torch.float).cuda()\n            len3=0\n\n        if (torch.sum(m2)+torch.sum(m3))!=0 :\n           res2 = torch.sum(m2)/max_pos+torch.sum(m3)/max_neg\n           #code.interact(local=dict(globals(), **locals()))\n        else:\n           res2 = torch.sum(m2)+torch.sum(m3)\n\n        res2 = torch.where(torch.isnan(res2), torch.zeros_like(res2), res2)\n\n        return res2\n    \nclass ROC_Star(nn.Module):\n    def __init__(self):\n        super().__init__()\n        \n    def forward(self, y_pred,_y_true,i):    #_epoch_true, epoch_pred):\n        return roc_star_loss( _y_true, y_pred, CFG.gamma, torch.from_numpy(CFG.last_epoch_true[i]).cuda(), torch.from_numpy(CFG.last_epoch_pred[i]).cuda())","metadata":{"execution":{"iopub.status.busy":"2021-09-11T07:50:09.202773Z","iopub.execute_input":"2021-09-11T07:50:09.203691Z","iopub.status.idle":"2021-09-11T07:50:09.232188Z","shell.execute_reply.started":"2021-09-11T07:50:09.203649Z","shell.execute_reply":"2021-09-11T07:50:09.231442Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train loop","metadata":{"papermill":{"duration":0.034375,"end_time":"2021-07-01T14:31:43.478039","exception":false,"start_time":"2021-07-01T14:31:43.443664","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# ====================================================\n# Train loop\n# ====================================================\ndef train_loop():\n    # ====================================================\n    # loader\n    # ====================================================\n    global FLAGS\n    fold, folds = FLAGS[\"fold\"], FLAGS[\"train\"]\n    device = xm.xla_device()\n\n    trn_idx = folds[folds['fold'] != fold].index\n    val_idx = folds[folds['fold'] == fold].index\n\n    train_folds = folds.loc[trn_idx].reset_index(drop=True)\n    valid_folds = folds.loc[val_idx].reset_index(drop=True)\n    valid_labels = valid_folds[CFG.target_col].values\n\n    train_dataset = TrainDataset(train_folds, transform=get_transforms(data='train'))\n    valid_dataset = TrainDataset(valid_folds, transform=get_transforms(data='train'))\n\n    train_sampler = torch.utils.data.distributed.DistributedSampler(\n        train_dataset,\n        num_replicas=xm.xrt_world_size(),\n        rank=xm.get_ordinal(),\n        shuffle=True)\n\n\n    valid_sampler = torch.utils.data.distributed.DistributedSampler(\n        valid_dataset,\n        num_replicas=xm.xrt_world_size(),\n        rank=xm.get_ordinal(),\n        shuffle=False,\n        )\n\n    train_loader = DataLoader(train_dataset,\n                              batch_size=CFG.batch_size,\n                              sampler=train_sampler,\n                              num_workers=CFG.num_workers, pin_memory=True, drop_last=True)\n    valid_loader = DataLoader(valid_dataset, \n                              batch_size=CFG.batch_size,\n                              sampler=valid_sampler,\n                              num_workers=CFG.num_workers, pin_memory=True, drop_last=False)\n    \n    # ====================================================\n    # scheduler \n    # ====================================================\n    def get_scheduler(optimizer):\n        if CFG.scheduler=='ReduceLROnPlateau':\n            scheduler = ReduceLROnPlateau(optimizer, mode='min', factor=CFG.factor, patience=CFG.patience, verbose=True, eps=CFG.eps)\n        elif CFG.scheduler=='CosineAnnealingLR':\n            scheduler = CosineAnnealingLR(optimizer, T_max=CFG.T_max, eta_min=CFG.min_lr, last_epoch=-1)\n        elif CFG.scheduler=='CosineAnnealingWarmRestarts':\n            scheduler = CosineAnnealingWarmRestarts(optimizer, T_0=CFG.T_0, T_mult=1, eta_min=CFG.min_lr, last_epoch=-1)\n        return scheduler\n\n    # ====================================================\n    # model & optimizer\n    # ====================================================\n    model = MX.to(device)\n\n    optimizer = Adam(model.parameters(), lr=CFG.lr*xm.xrt_world_size(), weight_decay=CFG.weight_decay, amsgrad=False)\n    scheduler = get_scheduler(optimizer)\n\n    \n    # ====================================================\n    # loop\n    # ====================================================\n    criterion = ROC_Star()\n    best_score = 0.\n    best_loss = np.inf\n    \n    for epoch in range(CFG.epochs):\n        xm.master_print(f\"Epoch :{epoch}\")\n        start_time = time.time()\n        \n        # train\n        para_loader = pl.ParallelLoader(train_loader, [device])\n        avg_loss = train_fn(fold, para_loader.per_device_loader(device), model, criterion, optimizer, epoch, scheduler, device)\n        xm.master_print(f\"Train Loss:{avg_loss:.4f}\")\n        # eval\n        para_loader = pl.ParallelLoader(valid_loader, [device])\n        avg_val_loss, preds, score = valid_fn(para_loader.per_device_loader(device), model, criterion, device)        \n        xm.master_print(f\"Val Loss:{avg_val_loss:.4f} Val Score:{score:.4f}\")\n        \n        if isinstance(scheduler, ReduceLROnPlateau):\n            scheduler.step(avg_val_loss)\n        elif isinstance(scheduler, CosineAnnealingLR):\n            scheduler.step()\n        elif isinstance(scheduler, CosineAnnealingWarmRestarts):\n            scheduler.step()\n\n        elapsed = time.time() - start_time\n        \n        xm.master_print(elapsed)\n        xm.rendezvous(\"epoch complete\")\n        if score > best_score:\n            best_score = score\n            xm.save(model.state_dict(), OUTPUT_DIR+f'{CFG.model_name}_fold{fold}_best_score.pth')\n        \n        if avg_val_loss < best_loss:\n            best_loss = avg_val_loss\n            xm.save(model.state_dict(), OUTPUT_DIR+f'{CFG.model_name}_fold{fold}_best_score.pth')\n    \n    xm.master_print('='*20)\n    xm.master_print(f'best_loss: {best_loss:.4f}')\n    xm.master_print(f'Score: {best_score:.4f}')\n    xm.master_print('='*20)\n\n    return best_score, best_loss","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","papermill":{"duration":0.053758,"end_time":"2021-07-01T14:31:43.566555","exception":false,"start_time":"2021-07-01T14:31:43.512797","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2021-09-11T07:50:09.23376Z","iopub.execute_input":"2021-09-11T07:50:09.234258Z","iopub.status.idle":"2021-09-11T07:50:09.259385Z","shell.execute_reply.started":"2021-09-11T07:50:09.234211Z","shell.execute_reply":"2021-09-11T07:50:09.25836Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def _mp_fn(rank, flags):\n    torch.set_default_tensor_type('torch.FloatTensor')\n    a = train_loop()","metadata":{"_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","papermill":{"duration":1704.25542,"end_time":"2021-07-01T15:00:07.939127","exception":false,"start_time":"2021-07-01T14:31:43.683707","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2021-09-11T07:50:09.260659Z","iopub.execute_input":"2021-09-11T07:50:09.260914Z","iopub.status.idle":"2021-09-11T07:50:09.280146Z","shell.execute_reply.started":"2021-09-11T07:50:09.260885Z","shell.execute_reply":"2021-09-11T07:50:09.278982Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for fold in range(CFG.n_fold):\n    if fold in CFG.trn_fold:\n        FLAGS={\"fold\": fold, \"train\": train}\n        xmp.spawn(_mp_fn, args=(FLAGS,), nprocs=8, start_method='fork')","metadata":{"execution":{"iopub.status.busy":"2021-09-11T07:50:09.281853Z","iopub.execute_input":"2021-09-11T07:50:09.282193Z","iopub.status.idle":"2021-09-11T07:54:28.961869Z","shell.execute_reply.started":"2021-09-11T07:50:09.282154Z","shell.execute_reply":"2021-09-11T07:54:28.960735Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}